diff --git a/.github/scripts/verify_otel_api_compatibility.py b/.github/scripts/verify_otel_api_compatibility.py index b0e88c7aa..2ff92fd69 100644 --- a/.github/scripts/verify_otel_api_compatibility.py +++ b/.github/scripts/verify_otel_api_compatibility.py @@ -142,11 +142,14 @@ def run_matrix(args: argparse.Namespace) -> int: for view in ("otel-invocation", "otel-execution"): name = f"{label}-api{version}-{view}" negative = label == "old-old" and version == "1.49.0" + incompatible_core = label == "old-new" case: dict[str, object] = {"name": name, "expected_negative_control": negative, + "expected_core_rejection": incompatible_core, "api": jar_facts(api), "context": jar_facts(context)} command = [args.java, "-cp", os.pathsep.join(map(str, [classes, core, plugin, *dependencies])), PROBE, str(core), str(plugin), str(api), str(context), view, - str(negative).lower(), str(version == "1.66.0").lower()] + str(negative).lower(), str(version == "1.66.0").lower(), + str(incompatible_core).lower()] try: log = output / f"{name}.log" execute(command, log, env=probe_environment(view), timeout=90) @@ -155,6 +158,8 @@ def run_matrix(args: argparse.Namespace) -> int: raise RuntimeError("Probe did not report successful completion") if negative and "NEGATIVE_CONTROL_REPRODUCED" not in contents: raise RuntimeError("Released negative control did not reproduce the reported failure") + if incompatible_core and "CORE_LIFECYCLE_REJECTION_CONFIRMED" not in contents: + raise RuntimeError("Old core must reject the new plugin before invocation startup") case["passed"] = True except RuntimeError as error: failures += 1 diff --git a/.github/workflows/otel-conformance-tests.yml b/.github/workflows/otel-conformance-tests.yml index 452db089d..bd7a4a9f6 100644 --- a/.github/workflows/otel-conformance-tests.yml +++ b/.github/workflows/otel-conformance-tests.yml @@ -68,8 +68,9 @@ jobs: actions: write contents: read id-token: write - uses: aws/aws-durable-execution-conformance-tests/.github/workflows/opentelemetry-orchestrator.yml@a66037abbbfa55fde97f714e30f0bc262edefd63 + uses: aws/aws-durable-execution-conformance-tests/.github/workflows/opentelemetry-orchestrator.yml@1768f32d61d958b974ce9e66748df72e23bc19b2 with: + case_count: 20 runs_on: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository && 'ubuntu-latest' || format('codebuild-github-actions-runner-{0}-{1}', github.run_id, github.run_attempt) }} language: java resource_prefix: j diff --git a/examples/src/test/java/software/amazon/lambda/durable/examples/CloudBasedIntegrationTest.java b/examples/src/test/java/software/amazon/lambda/durable/examples/CloudBasedIntegrationTest.java index 7490961a5..58b5cf8bd 100644 --- a/examples/src/test/java/software/amazon/lambda/durable/examples/CloudBasedIntegrationTest.java +++ b/examples/src/test/java/software/amazon/lambda/durable/examples/CloudBasedIntegrationTest.java @@ -21,6 +21,7 @@ import software.amazon.awssdk.regions.Region; import software.amazon.awssdk.services.lambda.LambdaClient; import software.amazon.awssdk.services.lambda.model.ErrorObject; +import software.amazon.awssdk.services.lambda.model.Event; import software.amazon.awssdk.services.lambda.model.OperationStatus; import software.amazon.awssdk.services.sts.StsClient; import software.amazon.lambda.durable.TypeToken; @@ -818,7 +819,20 @@ void testConcurrentWaitForConditionExample() { lambdaClient); var result = runner.run(new ConcurrentWaitForConditionExample.Input(3, 100, 50)); - assertEquals(ExecutionStatus.SUCCEEDED, result.getStatus()); + try { + assertEquals(ExecutionStatus.SUCCEEDED, result.getStatus()); + } catch (RuntimeException | AssertionError failure) { + try { + // Preserve the already-fetched service timeline without payloads, callback IDs, or checkpoint tokens. + var history = result.getHistoryEvents().stream() + .map(CloudBasedIntegrationTest::historyMetadata) + .toList(); + System.err.println("Concurrent wait-for-condition history: " + new JacksonSerDes().serialize(history)); + } catch (RuntimeException diagnosticsFailure) { + failure.addSuppressed(diagnosticsFailure); + } + throw failure; + } // Verify each operation finished with 3 attempts var allOperationsOutput = result.getResult(); @@ -840,6 +854,30 @@ void testConcurrentWaitForConditionExample() { } } + private static Map historyMetadata(Event event) { + var row = new HashMap(); + row.put("eventId", event.eventId()); + row.put("type", event.eventTypeAsString()); + row.put("operationId", event.id()); + row.put("parentId", event.parentId()); + row.put("name", event.name()); + row.put("at", String.valueOf(event.eventTimestamp())); + var retries = event.stepSucceededDetails() != null + ? event.stepSucceededDetails().retryDetails() + : event.stepFailedDetails() != null ? event.stepFailedDetails().retryDetails() : null; + if (retries != null) { + row.put("attempt", retries.currentAttempt()); + row.put("nextAttemptDelaySeconds", retries.nextAttemptDelaySeconds()); + } + var invocation = event.invocationCompletedDetails(); + if (invocation != null) { + row.put("invocationStart", String.valueOf(invocation.startTimestamp())); + row.put("invocationEnd", String.valueOf(invocation.endTimestamp())); + row.put("invocationFailed", invocation.error() != null); + } + return row; + } + @Test void testPluginExample() { var runner = diff --git a/examples/src/test/java/software/amazon/lambda/durable/examples/parallel/ParallelFailureToleranceExampleTest.java b/examples/src/test/java/software/amazon/lambda/durable/examples/parallel/ParallelFailureToleranceExampleTest.java index f1518e2ac..cb23be650 100644 --- a/examples/src/test/java/software/amazon/lambda/durable/examples/parallel/ParallelFailureToleranceExampleTest.java +++ b/examples/src/test/java/software/amazon/lambda/durable/examples/parallel/ParallelFailureToleranceExampleTest.java @@ -5,8 +5,27 @@ import static org.junit.jupiter.api.Assertions.*; 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.TimeUnit; +import java.util.stream.Collectors; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +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.TypeToken; +import software.amazon.lambda.durable.model.ConcurrencyCompletionStatus; import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.model.ParallelResult; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.OperationEndInfo; +import software.amazon.lambda.durable.plugin.UserFunctionEndInfo; +import software.amazon.lambda.durable.plugin.UserFunctionStartInfo; +import software.amazon.lambda.durable.serde.JacksonSerDes; import software.amazon.lambda.durable.testing.LocalDurableTestRunner; class ParallelFailureToleranceExampleTest { @@ -41,19 +60,104 @@ void succeedsWhenAllBranchesSucceed() { assertEquals(3, output.succeeded()); } - @Test - void failsWhenFailuresExceedTolerance() { + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void failsWhenFailuresExceedTolerance(boolean holdHealthyBranch) throws Exception { var handler = new ParallelFailureToleranceExample(); - var runner = LocalDurableTestRunner.create(ParallelFailureToleranceExample.Input.class, handler); + var probe = new CompletionProbe(holdHealthyBranch); + var config = DurableConfig.builder().withPlugins(probe.newPlugin()).build(); + var runner = LocalDurableTestRunner.create(ParallelFailureToleranceExample.Input.class, handler) + .withDurableConfig(config); + var caller = Executors.newSingleThreadExecutor(); + try { + var input = new ParallelFailureToleranceExample.Input(List.of("svc-a", "bad-svc-b", "bad-svc-c"), 1, 2); + var invocation = caller.submit(() -> runner.runUntilComplete(input)); + if (holdHealthyBranch) { + await(probe.parallelStored); + probe.releaseHealthy.countDown(); + } + var result = invocation.get(5, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.SUCCEEDED, result.getStatus()); + var output = result.getResult(ParallelFailureToleranceExample.Output.class); + assertEquals(2, output.failed()); - // 2 bad services, toleratedFailureCount=1 — second failure exceeds tolerance - var input = new ParallelFailureToleranceExample.Input(List.of("svc-a", "bad-svc-b", "bad-svc-c"), 1, 2); - var result = runner.runUntilComplete(input); + var operation = result.getOperation("call-services"); + assertEquals(OperationStatus.SUCCEEDED, operation.getStatus()); + var stored = new JacksonSerDes() + .deserialize(operation.getContextDetails().result(), TypeToken.get(ParallelResult.class)); + assertEquals(ConcurrencyCompletionStatus.FAILURE_TOLERANCE_EXCEEDED, stored.completionStatus()); + assertFalse(stored.completionStatus().isSucceeded()); + assertEquals(3, stored.size()); + assertEquals( + List.of(ParallelResult.Status.FAILED, ParallelResult.Status.FAILED), + stored.statuses().subList(1, 3)); + assertTrue(List.of(ParallelResult.Status.SUCCEEDED, ParallelResult.Status.SKIPPED) + .contains(stored.statuses().get(0))); + assertEquals(output.succeeded(), stored.succeeded()); + assertEquals(output.failed(), stored.failed()); + assertEquals(1 - output.succeeded(), stored.skipped()); + if (holdHealthyBranch) assertEquals(0, output.succeeded()); - assertEquals(ExecutionStatus.SUCCEEDED, result.getStatus()); + var completedCalls = result.getOperations().stream() + .filter(op -> op.getType() == OperationType.STEP && op.isCompleted()) + .collect(Collectors.toMap(op -> op.getName(), op -> probe.completedCalls.get(op.getName()))); + assertEquals(1, completedCalls.get("invoke-bad-svc-b")); + assertEquals(1, completedCalls.get("invoke-bad-svc-c")); + var replay = runner.run(input); + assertEquals(ExecutionStatus.SUCCEEDED, replay.getStatus()); + assertEquals(output, replay.getResult(ParallelFailureToleranceExample.Output.class)); + completedCalls.forEach((name, count) -> assertEquals( + count, + probe.completedCalls.get(name), + "Completed step bodies must not run again on replay: " + name)); + } finally { + probe.releaseHealthy.countDown(); + caller.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + } + } - var output = result.getResult(ParallelFailureToleranceExample.Output.class); - assertEquals(2, output.failed()); - assertEquals(1, output.succeeded()); + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(5, TimeUnit.SECONDS), "Controlled branch scheduling was not released"); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + + private static final class CompletionProbe { + final boolean holdHealthy; + final CountDownLatch healthyStarted = new CountDownLatch(1); + final CountDownLatch releaseHealthy = new CountDownLatch(1); + final CountDownLatch parallelStored = new CountDownLatch(1); + final Map completedCalls = new ConcurrentHashMap<>(); + + CompletionProbe(boolean holdHealthy) { + this.holdHealthy = holdHealthy; + } + + DurableExecutionPlugin newPlugin() { + return new DurableExecutionPlugin() { + @Override + public void onUserFunctionStart(UserFunctionStartInfo info) { + if (!holdHealthy || !"STEP".equals(info.type())) return; + if ("invoke-svc-a".equals(info.name())) { + healthyStarted.countDown(); + await(releaseHealthy); + } else if (info.name().startsWith("invoke-bad-")) await(healthyStarted); + } + + @Override + public void onUserFunctionEnd(UserFunctionEndInfo info) { + if ("STEP".equals(info.type())) completedCalls.merge(info.name(), 1, Integer::sum); + } + + @Override + public void onOperationEnd(OperationEndInfo info) { + if ("call-services".equals(info.name())) parallelStored.countDown(); + } + }; + } } } diff --git a/otel-plugin/README.md b/otel-plugin/README.md index 7a138f38c..9596ebf53 100644 --- a/otel-plugin/README.md +++ b/otel-plugin/README.md @@ -13,6 +13,79 @@ OpenTelemetry instrumentation plugin for the AWS Lambda Durable Execution SDK fo - **ADOT Java Agent Integration**: `new InvocationOtelPlugin()` late-binds the ADOT Java agent's global provider with no handler-side OpenTelemetry initialization - **Lambda Layer Discovery**: `DURABLE_EXECUTION_PLUGINS` loads either OTel plugin from a JAR under a layer's `java/lib` directory +## Root handler context + +User instrumentation in the root handler joins its canonical execution trace. A valid same-trace ambient Lambda +span stays current. If ambient context is absent or belongs to another trace, invocation view activates `Invocation`; +execution view activates the deterministic `Workflow` context, including its unsampled non-recording form. +Nested operations retain their existing contexts. + +The plugin activates the context in `onInvocationStart` and restores the previous context in `onInvocationEnd`. +The core calls both hooks on the root handler thread and waits for `onInvocationEnd` to finish before returning, +including when execution suspends or terminates. Handler `finally` blocks must finish before invocation-end hooks +can run; a blocked handler cleanup therefore also blocks the invocation response. Context restoration still runs +when span finalization or flushing fails. Ordinary exceptions and nonfatal linkage errors retain the existing +plugin-hook isolation behavior. Invocation-end hooks run in reverse registration order so nested scopes unwind +correctly; remaining end hooks run even when another hook raises an Error. The first unisolated Error propagates, +unless a later `VirtualMachineError` or `ThreadDeath` takes precedence over a non-JVM-fatal Error; other distinct +end-hook Errors are retained as suppressed failures. Scope cleanup preserves finalization failures using the same +JVM-fatal precedence, so an ordinary cleanup exception cannot hide an earlier Error. +Invocations without plugins also wait for handler cleanup before returning the selected suspension or retry outcome. +The manager selects suspension or termination before waking operation waiters, so a later returning or throwing +handler `finally` cannot replace that selected outcome. A handler outcome that was already selected remains primary. +SDK output preparation, including customer `SerDes` calls and durable large-result checkpointing, finishes before +terminal invocation-end notification. Failures in this preparation report `RETRYING` instead of ending the Workflow +span. If End also raises an unisolated Error, the preparation failure remains primary with the cleanup error +suppressed, unless cleanup introduces the first JVM-fatal error. An original JVM-fatal preparation error retains its +identity. Retryable control failures follow the same End-error combination: an ordinary End Error is suppressed +under the retry control, while a direct JVM-fatal End error retains its existing precedence with the control and +its cause preserved as diagnostics. This does not add arbitrary cause unwrapping to handler failure classification. +End describes the SDK outcome at that point, not acknowledgment of a response by the Lambda service. +Caller-side execution-manager cleanup, response-envelope encoding and output-stream writes follow End; runtime +response transport follows the handler return. Failures at those later boundaries still propagate, without a second +End dispatch or changing its already reported outcome. A JVM-fatal MDC-restoration failure after End (direct or +inside a standard transport wrapper) completes the caller's observation exceptionally with the original fatal +before that fatal escapes the worker, for both asynchronous and inline executors. The End notification remains +unchanged and is not repeated. An earlier JVM-fatal End/preparation failure remains primary; later restoration +failures are suppressed, with an identity guard when restoration throws the same fatal object. A restoration fatal takes precedence over an earlier non-JVM-fatal delivery/End failure, +which remains suppressed. Non-JVM-fatal restoration failures, including `AssertionError` and `LinkageError`, retain +the selected caller outcome. After an ordinary outer restoration failure, the SDK makes one guarded `MDC.clear()` +attempt before rethrowing the original worker failure. This also applies when an earlier End/preparation JVM fatal +is already primary: restoration and clear failures remain suppressed under that first fatal. A successful clear +prevents inheritable invocation state from reaching a replacement pool thread. With no earlier JVM fatal, a failing +clear is retained as a suppressed diagnostic, or its JVM fatal uses the same caller-settlement-before-worker-throw +rule. Clearing can discard ambient MDC that could not be +restored. If the adapter's clear also fails, clean replacement state is not guaranteed; the SDK does not promise +quarantine for an arbitrary caller-owned executor. Initialization and handler/body policies are unchanged. +Before invocation startup, a JVM-fatal error from MDC capture (direct or inside a standard transport wrapper) +completes the observation future exceptionally with that same fatal before escaping the handler worker. No start, +body, or end hook runs, and no durable `FAILED` response is produced for that fatal. Ordinary initialization errors +and handler/body failure classification retain their existing behavior. + +An SDK checkpoint continuation owned by an operation reports resumption/deserialization and dispatch failures +through retryable manager control before releasing its activity lease. The caller receives the original failure as +the cause of `UnrecoverableDurableExecutionException`, and started End hooks receive `RETRYING`; persisted operation +state is unchanged for the next invocation. A rejected user-worker dispatch rolls back its activity registration. +Ordinary failures during normal manager closing do not replace the selected outcome or stop unrelated operations. +Unowned ordinary helper failures retain their observation-only behavior. Direct `VirtualMachineError` or +`ThreadDeath` failures also settle their observation future after lease release and before escaping the coordinator +worker. Observation cancellation cannot skip the actual task or discard its operation failure. This boundary does +not reclassify handler/predicate failures or add general wrapped-fatal classification. + +SDK inspection of MDC-capture failures reads each visited standard transport cause once and detects identity cycles. +Cyclic, null, or unreadable leading `CompletionException` chains retain the original wrapper; ordinary initialization +still reports `FAILED` when its error response can be serialized. This provides cycle safety for finite cause graphs, +not a fixed depth or time limit. Arbitrary custom `getCause`, other `Throwable` accessors, and customer `SerDes` +behavior remain outside that guarantee. + +This plugin version requires the core's `DurableExecutor.supportsSameThreadInvocationHooks()` capability, introduced +in the 2.2.2 lifecycle contract (currently `2.2.2-SNAPSHOT`). Upgrade the core together with the plugin layer. Every +plugin constructor checks this capability before building a tracer provider or activating context. A core without +it, including released 2.2.1, is rejected with an explicit configuration error. Older cores may call +`onInvocationEnd` on another thread and are not supported with this plugin version. Existing plugin binaries remain +usable with the updated core; their invocation-end hooks now follow the same-thread, reverse-registration-order +contract described above. + ## Installation ```xml @@ -62,6 +135,13 @@ on it replaces the complete plugin list without reading `DURABLE_EXECUTION_PLUGI plugins from the copy. Use `DurableConfig.builder()` when creating a fresh configuration that should honor the current environment selection. +View exclusivity is declared with inherited `@ExclusivePluginGroup("durable-otel-view")` metadata. +Configuration reads this explicit opt-in annotation from the entire superclass chain; it does not call application methods +that happen to be named `getExclusiveGroup`. Existing subclasses retain their own methods while inheriting the bundled +view restriction. A subclass may add another group, but cannot replace a superclass's group; repeated group names in one +class hierarchy are checked once. +The exclusivity annotation does not change the provider registration API; the same-thread core requirement above still applies. + ## Quick Start using X-Ray/CloudWatch Tracing (ADOT Java Agent) 1. Add the ADOT Lambda Layer to your function @@ -224,14 +304,35 @@ Operation and attempt spans link to the Workflow span. `ExecutionOtelPlugin` rev ### Sampling -The plugin decides sampling once per invocation and applies that single decision to every durable span (Workflow, Invocation, operation, attempt), so the configured sampler is not re-invoked per span and the full decision — including `RECORD_ONLY` — is preserved. The decision follows this precedence, highest first: - -1. **Backend decision** — `Sampled=1` / `Sampled=0` in the propagated header is authoritative and always preserved, regardless of the configured sampler. -2. **Same-trace ambient span** — when the header carries no usable `Sampled` value but a valid ambient span (for example an auto-instrumentation Lambda handler span) is already on the execution's trace, the plugin follows that span's decision: sampled → sampled; unsampled but still recording → `RECORD_ONLY`; unsampled and not recording → dropped. -3. **Configured sampler (application-owned provider)** — when you pass a `SdkTracerProvider` to the plugin, its sampler is read directly and evaluated once with the trace ID, span name, and attributes. A trace-ID-ratio sampler therefore produces a stable decision across reinvocations (the trace ID is stable). -4. **Installed sampler (Java-agent path)** — when the agent owns the provider, it is behind a classloader boundary and its *effective* sampler (which another agent extension may have wrapped or replaced) cannot be reliably read at decision time. Rather than guess, the plugin **defers**: it installs a delegating sampler through the agent's autoconfiguration and lets that wrapper consult the agent's real sampler. The delegate's decision is honored in full — if your configured policy is `always_off`, a rate limiter, or a remote sampler (`xray`, `jaeger_remote`) that returns drop, the durable spans are dropped; they are **not** force-sampled. To avoid consuming a stateful or quota-based sampler once per span, the wrapper consults the delegate once per execution (keyed by trace ID) and reuses that decision for the execution's remaining durable spans within the invocation. - -For precise, provider-independent control, set an explicit `Sampled` value upstream (for example by enabling X-Ray active tracing) — that backend decision takes precedence over everything else. +The SDK-wide sampling guarantees below require `DurableSampler` itself to be the final installed sampler. The +[builder constructors](#configuration) and [ADOT extension setup](#1-adot-lambda-layer) install it automatically; +configure its delegate through these documented paths and retain `DurableSampler` as the final sampler. An +unrecognized outer wrapper is treated as a plain replacement. The SDK wrapper applies its sampling intent to every durable +span (Workflow, Invocation, operation and attempt), preserving the full result, including `RECORD_ONLY`, sampler attributes +and trace state. + +The supported SDK sampler follows this precedence, highest first: + +1. **Backend decision** — with `DurableSampler` installed, `Sampled=1` / `Sampled=0` in the propagated header is + authoritative, regardless of the wrapped delegate's policy. +2. **Same-trace ambient span** — without an explicit header decision, a valid ambient span already on the execution + trace supplies its decision: sampled → sampled; unsampled and recording → `RECORD_ONLY`; unsampled and + non-recording → dropped. +3. **Same-copy durable sampler** — the visible SDK sampler is evaluated against root context once per invocation + using the canonical trace ID, span name and attributes. Its complete result is carried in that loader's context. +4. **Foreign or opaque SDK sampler** — resolution is deferred to the installed wrapper's actual delegate. Its full + result is cached by execution ARN and canonical trace ID in a 256-entry LRU cache and reused while resident. + Eviction can cause another evaluation. The delegate still receives the canonical trace ID; drop and + `RECORD_ONLY` decisions, attributes and updated trace state are retained. + +A visible plain replacement sampler supplies its root policy for execution-ancestor flags, then keeps its normal +per-span sampling behavior. It receives no SDK sampling carrier and does not provide the SDK decision-reuse +or full-result override guarantees above. Opaque providers cannot expose a later sampler replacement to the +application, so retain the wrapper through the documented customization setup. + +For explicit execution-level control, keep `DurableSampler` as the final installed sampler and set `Sampled` upstream, +for example through X-Ray +active tracing. The SDK sampling wrapper preserves that decision. ## Span Attributes @@ -400,17 +501,10 @@ var otelPlugin = new InvocationOtelPlugin( ## Requirements - Java 17+ -- AWS Durable Execution SDK for Java 2.0.0+ +- AWS Durable Execution SDK for Java with same-thread invocation hooks (use the core shipped with this plugin release or newer) - OpenTelemetry SDK 1.65.0+ (only for custom TracerProvider path) - ADOT Lambda Layer `AWSOpenTelemetryDistroJava` (for the no-arg constructor path) ## License Apache-2.0 - -View exclusivity is declared with inherited `@ExclusivePluginGroup("durable-otel-view")` metadata. -Configuration reads this explicit opt-in annotation from the entire superclass chain; it does not call application methods -that happen to be named `getExclusiveGroup`. Existing subclasses retain their own methods while inheriting the bundled -view restriction. A subclass may add another group, but cannot replace a superclass's group; repeated group names in one -class hierarchy are checked once. -Older cores ignore the optional annotation and retain their prior behavior; no provider-version floor is raised. diff --git a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSampler.java b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSampler.java index 7ce304a54..78dc3d84e 100644 --- a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSampler.java +++ b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSampler.java @@ -44,13 +44,15 @@ final class DurableSampler implements Sampler { private static final int MAX_CACHED_DEFERRED_DECISIONS = 256; private final Sampler delegate; - // Caches the delegate's decision for a deferred durable execution, keyed by canonical trace ID, so a stateful or - // quota-based delegate is consulted once per execution rather than per span. Access-ordered LRU, size-capped and - // synchronized (contention is low: at most one miss per execution). - private final Map deferredDecisions = + // Caches the delegate's decision for a deferred durable execution, keyed by execution ARN and canonical trace ID, + // so a stateful or quota-based delegate is reused while that execution entry remains resident. + // Access-ordered LRU, size-capped and synchronized; eviction permits another delegate evaluation. + private record ExecutionKey(String traceId, String executionArn) {} + + private final Map deferredDecisions = Collections.synchronizedMap(new LinkedHashMap<>(16, 0.75f, true) { @Override - protected boolean removeEldestEntry(Map.Entry eldest) { + protected boolean removeEldestEntry(Map.Entry eldest) { return size() > MAX_CACHED_DEFERRED_DECISIONS; } }); @@ -113,7 +115,7 @@ public SamplingResult shouldSample( SpanKind spanKind, Attributes attributes, List parentLinks) { - var intent = DurableSamplingDecision.get(parentContext); + var intent = DurableSamplingDecision.consume(parentContext); if (intent == null) { // Not a durable span: the customer's sampler governs it unchanged. return delegate.shouldSample(parentContext, traceId, name, spanKind, attributes, parentLinks); @@ -126,8 +128,9 @@ public SamplingResult shouldSample( // Deferred (agent path, real sampler not reproducible here): evaluate the actual delegate once per execution // and reuse it, so an installed drop/rate-limit policy is honored and consulted only once. return deferredDecisions.computeIfAbsent( - intent.deferredTraceId(), - key -> delegate.shouldSample(Context.root(), key, name, spanKind, attributes, Collections.emptyList())); + new ExecutionKey(intent.deferredTraceId(), attributes.get(SpanAttributes.DURABLE_EXECUTION_ARN)), + key -> delegate.shouldSample( + Context.root(), key.traceId(), name, spanKind, attributes, Collections.emptyList())); } @Override diff --git a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSamplingDecision.java b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSamplingDecision.java index c43d8eb02..1d2de1217 100644 --- a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSamplingDecision.java +++ b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/DurableSamplingDecision.java @@ -6,6 +6,7 @@ import io.opentelemetry.context.ContextKey; import io.opentelemetry.sdk.trace.samplers.SamplingDecision; import io.opentelemetry.sdk.trace.samplers.SamplingResult; +import java.util.concurrent.atomic.AtomicReference; /** * Carries the durable execution's sampling intent to {@link DurableSampler} for one durable span. @@ -18,8 +19,8 @@ * preserving the full three-way decision (including {@code RECORD_ONLY}) and consulting no delegate; or *
  • a deferral marker (carrying the canonical trace ID) — used on the Java-agent path when the real sampler * (a remote/custom/file-only policy) cannot be reproduced here. {@link DurableSampler} then evaluates its actual - * delegate once per execution, caches the result by trace ID, and reuses it for the execution's remaining durable - * spans, so an installed drop/rate-limit policy is honored and consulted only once. + * delegate once per execution, caches the result by execution ARN and trace ID, and reuses it for the execution's + * remaining durable spans, so an installed drop/rate-limit policy is honored and consulted only once. * * *

    Two carriers, because the plugin runs across two class loaders. Under the documented ADOT setup the plugin @@ -28,28 +29,30 @@ * identity, so a key created in one loader is not equal to the key created in the other. To bridge this: * *

      - *
    1. Context key — used when both sides share a class loader (an application-owned provider). It preserves - * the full {@link SamplingResult}, including any attributes a custom sampler attached. + *
    2. Context key — used when both sides share a class loader (an application-owned provider). Agent-backed + * spans omit this carrier because another loader cannot consume its key. It preserves the full + * {@link SamplingResult}, including any attributes a custom sampler attached. *
    3. Thread-scoped system property — a cross-class-loader fallback modelled on * {@link DeterministicIdGenerator}'s scoped-ID bridge. The intent is published on the thread that creates the - * durable span for the synchronous duration of {@code startSpan()} (the sampler runs on that same thread), keyed - * by thread ID under a bootstrap-visible {@link System} property so both class loaders read the same value. A - * resolved decision bridges its {@link SamplingDecision} name (the three built-in decisions carry no attributes, - * so reconstructing them is faithful); a deferral bridges a sentinel plus the canonical trace ID. + * durable span until its sampler consumes the value (before synchronous span processors run), keyed by thread ID + * under a bootstrap-visible {@link System} property so both class loaders read the same value. A resolved + * decision bridges its {@link SamplingDecision} name (the three built-in decisions carry no attributes, so + * reconstructing them is faithful); a deferral bridges a sentinel plus the canonical trace ID. *
    * *

    The scope is opened immediately around each durable {@code startSpan()} call and closed right after, so the - * property never leaks beyond the span it applies to. Nothing is persisted across invocations; cross-invocation - * consistency comes from recomputing the intent from stable inputs, not from sharing state. + * property is consumed once by the sampler before callbacks can create unrelated spans. Nothing is persisted across + * invocations; cross-invocation consistency comes from recomputing the intent from stable inputs, not from sharing + * state. */ final class DurableSamplingDecision { /** * The durable sampling intent for a span: either a resolved {@link SamplingResult}, or a deferral to the agent-side - * sampler's own delegate keyed by the canonical trace ID. + * sampler's own delegate. The sampler uses the execution ARN already attached to the span with this trace ID. * * @param resolved the resolved decision, or null when deferring - * @param deferredTraceId the canonical trace ID to key the agent-side delegate cache on, or null when resolved + * @param deferredTraceId the canonical trace ID passed to the agent-side delegate, or null when resolved */ record Intent(SamplingResult resolved, String deferredTraceId) { static Intent resolved(SamplingResult result) { @@ -65,7 +68,7 @@ boolean isDeferred() { } } - private static final ContextKey KEY = + private static final ContextKey> KEY = ContextKey.named("software.amazon.lambda.durable.otel.durable-sampling-decision"); private static final String SCOPED_PROPERTY_PREFIX = "software.amazon.lambda.durable.otel.scopedSamplingDecision."; @@ -74,9 +77,9 @@ boolean isDeferred() { private DurableSamplingDecision() {} - /** Returns a context carrying the durable sampling intent (same-class-loader carrier), derived from the given. */ + /** Returns a context carrying a one-shot intent for one span, preserving the full SamplingResult. */ static Context store(Context context, Intent intent) { - return context.with(KEY, intent); + return context.with(KEY, new AtomicReference<>(intent)); } /** @@ -98,17 +101,30 @@ static Scope openScope(Intent intent) { } /** - * Returns the durable sampling intent for a span, or {@code null} when none is present. Prefers the full-fidelity - * context key (same class loader) and falls back to the thread-scoped system property (cross class loader). + * Inspects the unconsumed durable sampling intent for a span, or {@code null} when none is present. Prefers the + * full-fidelity context key (same class loader) and falls back to the thread-scoped system property (cross class + * loader). */ static Intent get(Context context) { - var fromContext = context.get(KEY); + var holder = context.get(KEY); + var fromContext = holder != null ? holder.get() : null; if (fromContext != null) { return fromContext; } return fromScopedProperty(); } + /** Consumes both carriers before callbacks can reuse the parent context or cross-loader bridge. */ + static Intent consume(Context context) { + try { + var holder = context.get(KEY); + var fromContext = holder != null ? holder.getAndSet(null) : null; + return fromContext != null ? fromContext : fromScopedProperty(); + } finally { + System.clearProperty(scopedProperty()); + } + } + private static String encode(Intent intent) { return intent.isDeferred() ? DEFERRED_PREFIX + intent.deferredTraceId() diff --git a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/ExecutionOtelPlugin.java b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/ExecutionOtelPlugin.java index 40faf0471..6f9a736fa 100644 --- a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/ExecutionOtelPlugin.java +++ b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/ExecutionOtelPlugin.java @@ -105,6 +105,8 @@ public class ExecutionOtelPlugin implements DurableExecutionPlugin { // Per-invocation state private volatile boolean tracingEnabled; private volatile Span invocationSpan; + // Opened and closed by the invocation hooks on the root handler thread. + private Scope handlerScope; private volatile String durableExecutionArn; // Trace ID and flags of the execution trace, published together as one snapshot so readers never pair a trace ID @@ -177,6 +179,7 @@ public ExecutionOtelPlugin() { * @param config the plugin configuration */ public ExecutionOtelPlugin(SdkTracerProviderBuilder tracerProviderBuilder, OtelPluginConfig config) { + OtelPluginSupport.requireSameThreadInvocationHooks(); this.idGenerator = DeterministicIdGenerator.installOn(tracerProviderBuilder); // Wrap the configured sampler so durable spans use the execution's single precomputed decision. DurableSampler.installOn(tracerProviderBuilder); @@ -198,6 +201,7 @@ public ExecutionOtelPlugin(SdkTracerProviderBuilder tracerProviderBuilder, OtelP * @param config the plugin configuration */ public ExecutionOtelPlugin(OtelPluginConfig config) { + OtelPluginSupport.requireSameThreadInvocationHooks(); this.contextExtractor = config.contextExtractor(); this.enableMdc = config.enableMdc(); this.workflowSpanName = config.workflowSpanName(); @@ -271,10 +275,27 @@ public void onInvocationStart(InvocationInfo info) { invocationSpan.getSpanContext().getTraceId()); } tracingEnabled = true; + handlerScope = activateHandlerContext(); + } + + private Scope activateHandlerContext() { + var trace = executionTrace; + if (!tracingEnabled || trace == null) return null; + var ambient = Span.current().getSpanContext(); + // Preserve a compatible ambient Lambda span. An absent or unrelated ambient span must not leave + // handler instrumentation outside the durable execution's canonical trace. + if (ambient.isValid() && trace.traceId().equals(ambient.getTraceId())) return Scope.noop(); + return Span.wrap(workflowSpanContext).makeCurrent(); } @Override public void onInvocationEnd(InvocationEndInfo info) { + var scope = handlerScope; + handlerScope = null; + OtelPluginSupport.runInvocationEnd(scope, () -> endInvocation(info)); + } + + private void endInvocation(InvocationEndInfo info) { if (!tracingEnabled) { return; } @@ -634,7 +655,9 @@ private Context resolveParentContext(String parentId) { */ private Context withDurableDecision(Context context) { var intent = samplingIntent; - return intent != null ? DurableSamplingDecision.store(context, intent) : context; + return intent != null && OtelPluginSupport.usesLocalDurableSampler(sdkTracerProvider) + ? DurableSamplingDecision.store(context, intent) + : context; } private TraceFlags effectiveTraceFlags() { @@ -659,7 +682,7 @@ private TraceState effectiveTraceState() { */ private Span startDurableSpan(SpanBuilder spanBuilder) { var intent = samplingIntent; - if (intent == null) { + if (intent == null || !OtelPluginSupport.usesDurableSamplingBridge(sdkTracerProvider)) { return spanBuilder.startSpan(); } try (var ignored = DurableSamplingDecision.openScope(intent)) { @@ -670,7 +693,7 @@ private Span startDurableSpan(SpanBuilder spanBuilder) { /** Starts a durable span with a forced span ID, publishing the sampling intent as in {@link #startDurableSpan}. */ private Span startDurableSpan(SpanBuilder spanBuilder, String traceId, String spanId) { var intent = samplingIntent; - if (intent == null) { + if (intent == null || !OtelPluginSupport.usesDurableSamplingBridge(sdkTracerProvider)) { return idGenerator.startSpan(spanBuilder, traceId, spanId); } try (var ignored = DurableSamplingDecision.openScope(intent)) { diff --git a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/InvocationOtelPlugin.java b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/InvocationOtelPlugin.java index 2928feea8..4e2e3cd70 100644 --- a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/InvocationOtelPlugin.java +++ b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/InvocationOtelPlugin.java @@ -92,6 +92,8 @@ public class InvocationOtelPlugin implements DurableExecutionPlugin { // Per-invocation state private volatile boolean tracingEnabled; private volatile Span invocationSpan; + // Opened and closed by the invocation hooks on the root handler thread. + private Scope handlerScope; private volatile String durableExecutionArn; // Trace ID and flags of the execution trace, published together as one snapshot so readers never pair a trace ID // with mismatched flags. @@ -170,6 +172,7 @@ public InvocationOtelPlugin() { * @param config the plugin configuration */ public InvocationOtelPlugin(SdkTracerProviderBuilder tracerProviderBuilder, OtelPluginConfig config) { + OtelPluginSupport.requireSameThreadInvocationHooks(); this.idGenerator = DeterministicIdGenerator.installOn(tracerProviderBuilder); // Wrap the configured sampler so durable spans use the execution's single precomputed decision. DurableSampler.installOn(tracerProviderBuilder); @@ -191,6 +194,7 @@ public InvocationOtelPlugin(SdkTracerProviderBuilder tracerProviderBuilder, Otel * @param config the plugin configuration */ public InvocationOtelPlugin(OtelPluginConfig config) { + OtelPluginSupport.requireSameThreadInvocationHooks(); this.contextExtractor = config.contextExtractor(); this.enableMdc = config.enableMdc(); this.workflowSpanName = config.workflowSpanName(); @@ -268,10 +272,27 @@ public void onInvocationStart(InvocationInfo info) { invocationSpan.getSpanContext().getTraceId()); } tracingEnabled = true; + handlerScope = activateHandlerContext(); + } + + private Scope activateHandlerContext() { + var trace = executionTrace; + if (!tracingEnabled || trace == null) return null; + var ambient = Span.current().getSpanContext(); + // Preserve a compatible ambient Lambda span. An absent or unrelated ambient span must not leave + // handler instrumentation outside the durable execution's canonical trace. + if (ambient.isValid() && trace.traceId().equals(ambient.getTraceId())) return Scope.noop(); + return invocationSpan.makeCurrent(); } @Override public void onInvocationEnd(InvocationEndInfo info) { + var scope = handlerScope; + handlerScope = null; + OtelPluginSupport.runInvocationEnd(scope, () -> endInvocation(info)); + } + + private void endInvocation(InvocationEndInfo info) { if (!tracingEnabled) { return; } @@ -633,6 +654,11 @@ private Context invocationParentContext(ExecutionTraceContext execCtx, String ca private Context resolveParentContext(String parentId) { if (parentId != null) { + var parentSpan = operationSpans.get(parentId); + if (parentSpan != null) { + // Retain the provider's live span, including its clock, while this parent is open. + return withDurableDecision(Context.current().with(parentSpan)); + } var parentSpanContext = operationContexts.get(parentId); if (parentSpanContext != null) { return withDurableDecision(Context.current().with(Span.wrap(parentSpanContext))); @@ -652,7 +678,9 @@ private Context resolveParentContext(String parentId) { */ private Context withDurableDecision(Context context) { var intent = samplingIntent; - return intent != null ? DurableSamplingDecision.store(context, intent) : context; + return intent != null && OtelPluginSupport.usesLocalDurableSampler(sdkTracerProvider) + ? DurableSamplingDecision.store(context, intent) + : context; } /** @@ -663,7 +691,7 @@ private Context withDurableDecision(Context context) { */ private Span startDurableSpan(SpanBuilder spanBuilder) { var intent = samplingIntent; - if (intent == null) { + if (intent == null || !OtelPluginSupport.usesDurableSamplingBridge(sdkTracerProvider)) { return spanBuilder.startSpan(); } try (var ignored = DurableSamplingDecision.openScope(intent)) { @@ -674,7 +702,7 @@ private Span startDurableSpan(SpanBuilder spanBuilder) { /** Starts a durable span with a forced span ID, publishing the sampling intent as in {@link #startDurableSpan}. */ private Span startDurableSpan(SpanBuilder spanBuilder, String traceId, String spanId) { var intent = samplingIntent; - if (intent == null) { + if (intent == null || !OtelPluginSupport.usesDurableSamplingBridge(sdkTracerProvider)) { return idGenerator.startSpan(spanBuilder, traceId, spanId); } try (var ignored = DurableSamplingDecision.openScope(intent)) { diff --git a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/OtelPluginSupport.java b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/OtelPluginSupport.java index dea21b447..e3af17bab 100644 --- a/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/OtelPluginSupport.java +++ b/otel-plugin/src/main/java/software/amazon/lambda/durable/otel/OtelPluginSupport.java @@ -9,6 +9,7 @@ import io.opentelemetry.api.trace.Tracer; import io.opentelemetry.api.trace.TracerProvider; import io.opentelemetry.context.Context; +import io.opentelemetry.context.Scope; import io.opentelemetry.sdk.trace.SdkTracerProvider; import io.opentelemetry.sdk.trace.samplers.SamplingResult; import java.nio.file.Files; @@ -16,6 +17,7 @@ import java.util.Collections; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import software.amazon.lambda.durable.execution.DurableExecutor; /** Shared utilities for OTel plugin default constructor support (ADOT Java agent SPI path). */ final class OtelPluginSupport { @@ -24,6 +26,59 @@ final class OtelPluginSupport { private OtelPluginSupport() {} + /** Only the same DurableSampler class copy can consume this loader's parent-context holder. */ + static boolean usesLocalDurableSampler(SdkTracerProvider provider) { + return provider != null && provider.getSampler() instanceof DurableSampler; + } + + /** A visible replacement sampler cannot consume our bridge; an opaque agent provider may still need it. */ + static boolean usesDurableSamplingBridge(SdkTracerProvider provider) { + // Class-name matching is only for the cross-loader wire carrier, never for the context-key ownership check. + return provider == null || provider.getSampler().getClass().getName().equals(DurableSampler.class.getName()); + } + + /** Closes the owning thread's scope without hiding a finalization error or a later JVM-fatal cleanup failure. */ + static void runInvocationEnd(Scope scope, Runnable end) { + Throwable primary = null; + try { + end.run(); + } catch (Throwable failure) { + primary = failure; + throw failure; + } finally { + try { + if (scope != null) scope.close(); + } catch (Throwable cleanup) { + if (primary == null) throw cleanup; + if (primary != cleanup) { + if (endFailurePriority(cleanup) > endFailurePriority(primary)) { + cleanup.addSuppressed(primary); + throw cleanup; + } + primary.addSuppressed(cleanup); + } + } + } + } + + @SuppressWarnings("removal") + private static int endFailurePriority(Throwable failure) { + if (failure instanceof VirtualMachineError || failure instanceof ThreadDeath) return 3; + if (failure instanceof Error && !(failure instanceof LinkageError)) return 2; + return 1; + } + + /** Rejects an older core before a plugin builds a provider or activates any thread-local context. */ + static void requireSameThreadInvocationHooks() { + var message = "This OpenTelemetry plugin requires same-thread invocation hooks from the Durable Execution" + + " core (2.2.2 lifecycle capability). Upgrade the core together with the plugin layer."; + try { + if (!DurableExecutor.supportsSameThreadInvocationHooks()) throw new IllegalStateException(message); + } catch (NoSuchMethodError missingCapability) { + throw new IllegalStateException(message, missingCapability); + } + } + /** Creates a new DeterministicIdGenerator for the application-side state bridge. */ static DeterministicIdGenerator createDefaultIdGenerator() { return new DeterministicIdGenerator(); @@ -56,7 +111,7 @@ static DeterministicIdGenerator createDefaultIdGenerator() { * pipeline finally installs, and another extension's customizer can wrap or replace a recognized configured * sampler, so a reconstruction could disagree with the real delegate. Deferring routes the decision to the * agent-installed {@link DurableSampler}, which consults its actual delegate once per execution, caches the - * result by trace ID, and reuses it for the execution's remaining durable spans (see + * result by execution ARN and trace ID, and reuses it for the execution's remaining durable spans (see * {@link DurableSampler#shouldSample}). The delegate's decision is honored in full — including a * {@code DROP}/rate-limited outcome — so durable spans are not force-sampled. * @@ -107,8 +162,11 @@ static SamplingResult resolveSamplingResult( } return ambientSpan.isRecording() ? SamplingResult.recordOnly() : SamplingResult.drop(); } - // 3. An application-owned provider exposes the real sampler: evaluate it once, preserving its full result. - if (sdkTracerProvider != null) { + // 3. Resolve a local durable sampler's full result, or a visible replacement's root policy for ancestor flags. + // A foreign DurableSampler still resolves in its own loader to retain full attributes and trace state. + // A plain replacement receives no carrier; the provider keeps its normal per-span sampling behavior. + if (sdkTracerProvider != null + && (usesLocalDurableSampler(sdkTracerProvider) || !usesDurableSamplingBridge(sdkTracerProvider))) { return sdkTracerProvider .getSampler() .shouldSample( diff --git a/otel-plugin/src/test/compatibility/b1/InstalledApiProbe.java b/otel-plugin/src/test/compatibility/b1/InstalledApiProbe.java index 9d31b1d58..26721887b 100644 --- a/otel-plugin/src/test/compatibility/b1/InstalledApiProbe.java +++ b/otel-plugin/src/test/compatibility/b1/InstalledApiProbe.java @@ -2,8 +2,8 @@ // SPDX-License-Identifier: Apache-2.0 package software.amazon.lambda.durable.otel; -import ch.qos.logback.classic.Logger; import ch.qos.logback.classic.Level; +import ch.qos.logback.classic.Logger; import ch.qos.logback.classic.spi.ILoggingEvent; import ch.qos.logback.core.read.ListAppender; import io.opentelemetry.api.GlobalOpenTelemetry; @@ -17,6 +17,7 @@ import java.time.Duration; import java.util.List; import java.util.ServiceLoader; +import java.util.concurrent.Callable; import java.util.concurrent.atomic.AtomicInteger; import org.slf4j.LoggerFactory; import software.amazon.lambda.durable.model.ExecutionStatus; @@ -43,7 +44,8 @@ public static void main(String[] args) throws Exception { logs.start(); root.addAppender(logs); try { - exercise(view, Boolean.parseBoolean(args[5]), Boolean.parseBoolean(args[6]), logs); + if (Boolean.parseBoolean(args[7])) assertOldCoreRejected(view); + else exercise(view, Boolean.parseBoolean(args[5]), Boolean.parseBoolean(args[6]), logs); System.out.println("COMPAT_PASS " + view + " negative=" + args[5] + " api=" + args[2] + " core=" + args[0] + " plugin=" + args[1]); } finally { @@ -66,6 +68,44 @@ private static void verifyArtifacts(String[] args) throws Exception { checkSource(provider.getPluginType(), Path.of(args[1])); } + private static void assertOldCoreRejected(String view) throws Exception { + var provider = ServiceLoader.load(DurableExecutionPluginProvider.class).stream() + .map(ServiceLoader.Provider::get).filter(value -> value.getName().equals(view)) + .findFirst().orElseThrow(); + expectCoreRejection(provider::createPlugin); + var constructors = provider.getPluginType().getConstructors(); + check(constructors.length == 4, "exercise every published plugin constructor"); + for (var constructor : constructors) { + var types = constructor.getParameterTypes(); + var parameters = new Object[types.length]; + for (int i = 0; i < types.length; i++) { + parameters[i] = types[i].getName().endsWith("OtelPluginConfig") + ? types[i].getMethod("defaults").invoke(null) : SdkTracerProvider.builder(); + } + expectCoreRejection(() -> constructor.newInstance(parameters)); + } + var handlerCalls = new AtomicInteger(); + expectCoreRejection(() -> createRunner(handlerCalls, new AtomicInteger())); + check(handlerCalls.get() == 0 && HEALTHY_STARTS.get() == 0 && HEALTHY_ENDS.get() == 0, + "configuration rejection must precede handler execution and lifecycle hooks"); + System.out.println("CORE_LIFECYCLE_REJECTION_CONFIRMED " + view); + } + + private static void expectCoreRejection(Callable construction) { + Throwable rejected = null; + try { + construction.call(); + } catch (Throwable failure) { + rejected = failure; + } + check(rejected != null, "new plugin silently accepted the released core without its lifecycle guarantee"); + for (var failure = rejected; failure != null; failure = failure.getCause()) { + if (failure instanceof IllegalStateException && failure.getMessage() != null + && failure.getMessage().contains("requires same-thread invocation hooks")) return; + } + throw new AssertionError("missing explicit core compatibility diagnostic", rejected); + } + private static void exercise(String view, boolean negative, boolean compatible, ListAppender logs) { GlobalOpenTelemetry.resetForTest(); // Reproduce #763's documented visible-API skew, not a claim of a deployed agent test. diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplerTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplerTest.java index 6fea5784e..f6126492e 100644 --- a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplerTest.java +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplerTest.java @@ -62,6 +62,32 @@ void tearDown() { OtelPluginAutoConfigurationState.resetInstalledForTest(); } + @Test + void consumingContextIntentAlsoClearsCrossLoaderFallback() { + var delegate = new CountingSampler(Sampler.alwaysOff()); + var sampler = DurableSampler.wrap(delegate); + var parent = DurableSamplingDecision.store( + Context.root(), DurableSamplingDecision.Intent.resolved(SamplingResult.recordOnly())); + try (var ignored = DurableSamplingDecision.openScope( + DurableSamplingDecision.Intent.resolved(SamplingResult.recordAndSample()))) { + assertEquals( + SamplingDecision.RECORD_ONLY, + sampler.shouldSample(parent, TRACE_ID, "durable", SpanKind.INTERNAL, Attributes.empty(), List.of()) + .getDecision()); + assertEquals( + SamplingDecision.DROP, + sampler.shouldSample( + Context.root(), + TRACE_ID, + "callback", + SpanKind.INTERNAL, + Attributes.empty(), + List.of()) + .getDecision()); + assertEquals(1, delegate.count(), "The unrelated callback must use its own sampler"); + } + } + // ─── Unit tests for the wrapper ────────────────────────────────────── @Test @@ -124,9 +150,10 @@ void deferredIntent_evaluatesDelegateOncePerExecution() { // A stateful/quota delegate must be consulted once per execution (trace ID), not per durable span. var delegate = new CountingSampler(Sampler.alwaysOn()); var sampler = DurableSampler.wrap(delegate); - var parent = DurableSamplingDecision.store(Context.root(), DurableSamplingDecision.Intent.deferred(TRACE_ID)); - for (var i = 0; i < 4; i++) { + // Each SDK-owned span receives its own one-shot carrier, while the resolved decision stays execution-wide. + var parent = + DurableSamplingDecision.store(Context.root(), DurableSamplingDecision.Intent.deferred(TRACE_ID)); sampler.shouldSample(parent, TRACE_ID, "op" + i, SpanKind.INTERNAL, Attributes.empty(), List.of()); } diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplingDecisionClassLoaderTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplingDecisionClassLoaderTest.java index 121f4c8fb..59b5cb32d 100644 --- a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplingDecisionClassLoaderTest.java +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/DurableSamplingDecisionClassLoaderTest.java @@ -7,13 +7,37 @@ import static org.junit.jupiter.api.Assertions.assertNotSame; import static org.junit.jupiter.api.Assertions.assertNull; +import io.opentelemetry.api.GlobalOpenTelemetry; +import io.opentelemetry.api.OpenTelemetry; +import io.opentelemetry.api.common.AttributeKey; +import io.opentelemetry.api.common.Attributes; +import io.opentelemetry.api.trace.SpanKind; +import io.opentelemetry.api.trace.TraceState; +import io.opentelemetry.api.trace.Tracer; +import io.opentelemetry.api.trace.TracerProvider; import io.opentelemetry.context.Context; +import io.opentelemetry.context.propagation.ContextPropagators; +import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter; +import io.opentelemetry.sdk.trace.IdGenerator; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.data.LinkData; +import io.opentelemetry.sdk.trace.export.SimpleSpanProcessor; +import io.opentelemetry.sdk.trace.samplers.Sampler; import io.opentelemetry.sdk.trace.samplers.SamplingDecision; import io.opentelemetry.sdk.trace.samplers.SamplingResult; import java.net.URL; import java.net.URLClassLoader; +import java.time.Instant; +import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; /** * Verifies the durable sampling decision crosses the application/Java-agent class-loader boundary. @@ -26,6 +50,129 @@ * then asserts the thread-scoped system-property bridge carries the decision from one loader to the other. */ class DurableSamplingDecisionClassLoaderTest { + @ParameterizedTest + @CsvSource({ + "InvocationOtelPlugin,local,false", "InvocationOtelPlugin,local,true", + "InvocationOtelPlugin,foreign,false", "InvocationOtelPlugin,foreign,true", + "InvocationOtelPlugin,opaque,false", "InvocationOtelPlugin,opaque,true", + "ExecutionOtelPlugin,local,false", "ExecutionOtelPlugin,local,true", + "ExecutionOtelPlugin,foreign,false", "ExecutionOtelPlugin,foreign,true", + "ExecutionOtelPlugin,opaque,false", "ExecutionOtelPlugin,opaque,true" + }) + void customSamplerMetadataSurvivesProviderOwnershipBoundaries( + String pluginName, String topology, boolean sharedTraceExecutions) throws Exception { + var previousHeader = System.getProperty("com.amazonaws.xray.traceHeader"); + GlobalOpenTelemetry.resetForTest(); + OtelPluginAutoConfigurationState.markInstalled(); + System.setProperty("com.amazonaws.xray.traceHeader", "Root=1-6955b900-123456789012345678901234"); + var evaluations = new AtomicInteger(); + var key = AttributeKey.stringKey("sampler.extra"); + var delegate = new Sampler() { + public SamplingResult shouldSample( + Context parent, + String traceId, + String name, + SpanKind kind, + Attributes attributes, + List links) { + evaluations.incrementAndGet(); + var second = + attributes.get(SpanAttributes.DURABLE_EXECUTION_ARN).contains("/second/"); + var metadata = sharedTraceExecutions ? (second ? "second" : "first") : "kept"; + return new SamplingResult() { + public SamplingDecision getDecision() { + return sharedTraceExecutions && !second + ? SamplingDecision.RECORD_ONLY + : SamplingDecision.RECORD_AND_SAMPLE; + } + + public Attributes getAttributes() { + return Attributes.of(key, metadata); + } + + public TraceState getUpdatedTraceState(TraceState parentState) { + return parentState.toBuilder().put("vendor", metadata).build(); + } + }; + } + + public String getDescription() { + return "custom-metadata"; + } + }; + try (var appLoader = pluginClassLoader(); + var agentLoader = pluginClassLoader(); + var exporter = InMemorySpanExporter.create()) { + var loader = topology.equals("local") ? appLoader : agentLoader; + var samplerType = Class.forName(DurableSampler.class.getName(), true, loader); + var wrap = samplerType.getDeclaredMethod("wrap", Sampler.class); + wrap.setAccessible(true); + var idType = Class.forName(DeterministicIdGenerator.class.getName(), true, loader); + try (var provider = SdkTracerProvider.builder() + .setSampler((Sampler) wrap.invoke(null, delegate)) + .setIdGenerator((IdGenerator) idType.getConstructor().newInstance()) + .addSpanProcessor(SimpleSpanProcessor.builder(exporter) + .setExportUnsampledSpans(true) + .build()) + .build()) { + var hidden = new TracerProvider() { + public Tracer get(String name) { + return provider.get(name); + } + + public Tracer get(String name, String version) { + return provider.get(name, version); + } + }; + GlobalOpenTelemetry.set(new OpenTelemetry() { + public TracerProvider getTracerProvider() { + return topology.equals("opaque") ? hidden : provider; + } + + public ContextPropagators getPropagators() { + return ContextPropagators.noop(); + } + }); + var plugin = (DurableExecutionPlugin) + Class.forName("software.amazon.lambda.durable.otel." + pluginName, true, appLoader) + .getConstructor() + .newInstance(); + var executionCount = sharedTraceExecutions ? 2 : 1; + for (var index = 0; index < executionCount; index++) { + var executionName = index == 0 ? "first" : "second"; + var arn = "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/" + executionName + + "/id"; + for (var first : new boolean[] {true, false}) { + plugin.onInvocationStart(new InvocationInfo("request", arn, first, Instant.ofEpochSecond(10))); + plugin.onInvocationEnd( + new InvocationEndInfo("request", arn, first, InvocationStatus.PENDING, null)); + } + } + var spans = exporter.getFinishedSpanItems().stream() + .filter(s -> s.getName().equals("Invocation")) + .toList(); + assertEquals(executionCount * 2, spans.size()); + for (var span : spans) { + var second = span.getAttributes() + .get(SpanAttributes.DURABLE_EXECUTION_ARN) + .contains("/second/"); + var expected = sharedTraceExecutions ? (second ? "second" : "first") : "kept"; + assertEquals(expected, span.getAttributes().get(key), span.getName()); + assertEquals(expected, span.getSpanContext().getTraceState().get("vendor"), span.getName()); + assertEquals( + !sharedTraceExecutions || second, + span.getSpanContext().isSampled()); + assertEquals("6955b900123456789012345678901234", span.getTraceId()); + } + assertEquals(executionCount * (topology.equals("local") ? 2 : 1), evaluations.get()); + } + } finally { + GlobalOpenTelemetry.resetForTest(); + OtelPluginAutoConfigurationState.resetInstalledForTest(); + if (previousHeader == null) System.clearProperty("com.amazonaws.xray.traceHeader"); + else System.setProperty("com.amazonaws.xray.traceHeader", previousHeader); + } + } @AfterEach void clearBridge() { diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/ExecutionOtelPluginTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/ExecutionOtelPluginTest.java index 8b84efc10..875ae6c57 100644 --- a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/ExecutionOtelPluginTest.java +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/ExecutionOtelPluginTest.java @@ -280,7 +280,8 @@ void workflowAndInvocationSpans_shareExecutionTrace_withoutAmbientContext() { void invocationStart_joinsAmbientTrace_whenAmbientIsOnExecutionTrace() { // Drive an invocation to learn the canonical execution trace ID, then start a fresh invocation with an ambient // span on that same trace: the Invocation span joins the ambient span directly. - plugin.onInvocationStart(new InvocationInfo("req-0", ARN, true, Instant.now())); + var executionStart = Instant.parse("2026-01-01T00:00:00Z"); + plugin.onInvocationStart(new InvocationInfo("req-0", ARN, true, executionStart)); plugin.onInvocationEnd(new InvocationEndInfo("req-0", ARN, true, InvocationStatus.SUCCEEDED, null)); var canonicalTraceId = spanByName(spanExporter.getFinishedSpanItems(), "Workflow").getTraceId(); @@ -290,9 +291,9 @@ void invocationStart_joinsAmbientTrace_whenAmbientIsOnExecutionTrace() { var ambient = SpanContext.create(canonicalTraceId, ambientSpanId, TraceFlags.getSampled(), TraceState.getDefault()); try (var ignored = Span.wrap(ambient).makeCurrent()) { - plugin.onInvocationStart(new InvocationInfo("req-1", ARN, false, Instant.now())); + plugin.onInvocationStart(new InvocationInfo("req-1", ARN, false, executionStart)); + plugin.onInvocationEnd(new InvocationEndInfo("req-1", ARN, false, InvocationStatus.SUCCEEDED, null)); } - plugin.onInvocationEnd(new InvocationEndInfo("req-1", ARN, false, InvocationStatus.SUCCEEDED, null)); var invocationSpan = spanByName(spanExporter.getFinishedSpanItems(), "Invocation"); assertEquals(canonicalTraceId, invocationSpan.getTraceId()); @@ -311,8 +312,8 @@ void invocationStart_staysOnExecutionTrace_withoutLinkingAmbientSpan() { SpanContext.create(ambientTraceId, ambientSpanId, TraceFlags.getSampled(), TraceState.getDefault()); try (var ignored = Span.wrap(ambient).makeCurrent()) { plugin.onInvocationStart(new InvocationInfo("req-1", ARN, true, Instant.now())); + plugin.onInvocationEnd(new InvocationEndInfo("req-1", ARN, true, InvocationStatus.SUCCEEDED, null)); } - plugin.onInvocationEnd(new InvocationEndInfo("req-1", ARN, true, InvocationStatus.SUCCEEDED, null)); var spans = spanExporter.getFinishedSpanItems(); var workflowSpan = spanByName(spans, "Workflow"); @@ -349,8 +350,9 @@ void contextExtractor_isInvokedEveryInvocation_evenWithAmbientSpan_andBackendCon TraceState.getDefault()); try (var ignored = Span.wrap(ambient).makeCurrent()) { extractorPlugin.onInvocationStart(new InvocationInfo("req-1", ARN, true, Instant.now())); + extractorPlugin.onInvocationEnd( + new InvocationEndInfo("req-1", ARN, true, InvocationStatus.SUCCEEDED, null)); } - extractorPlugin.onInvocationEnd(new InvocationEndInfo("req-1", ARN, true, InvocationStatus.SUCCEEDED, null)); assertEquals(1, extractCalls.get(), "Extractor is invoked even when a valid ambient span is active"); var spans = exporter.getFinishedSpanItems(); @@ -377,16 +379,16 @@ void executionTrace_isStableAcrossReinvocations_withDifferentAmbientTraces() { try (var ignored = Span.wrap(ambientA).makeCurrent()) { plugin.onInvocationStart(new InvocationInfo("req-1", ARN, true, startTime)); + plugin.onInvocationEnd(new InvocationEndInfo("req-1", ARN, true, InvocationStatus.PENDING, null)); } - plugin.onInvocationEnd(new InvocationEndInfo("req-1", ARN, true, InvocationStatus.PENDING, null)); var firstInvocationTrace = spanByName(spanExporter.getFinishedSpanItems(), "Invocation").getTraceId(); spanExporter.reset(); try (var ignored = Span.wrap(ambientB).makeCurrent()) { plugin.onInvocationStart(new InvocationInfo("req-2", ARN, false, startTime)); + plugin.onInvocationEnd(new InvocationEndInfo("req-2", ARN, false, InvocationStatus.SUCCEEDED, null)); } - plugin.onInvocationEnd(new InvocationEndInfo("req-2", ARN, false, InvocationStatus.SUCCEEDED, null)); var secondInvocationTrace = spanByName(spanExporter.getFinishedSpanItems(), "Invocation").getTraceId(); diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/HandlerContextIntegrationTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/HandlerContextIntegrationTest.java new file mode 100644 index 000000000..faa242bc0 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/HandlerContextIntegrationTest.java @@ -0,0 +1,290 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import io.opentelemetry.api.trace.Span; +import io.opentelemetry.api.trace.SpanContext; +import io.opentelemetry.api.trace.TraceFlags; +import io.opentelemetry.api.trace.TraceState; +import io.opentelemetry.context.Context; +import io.opentelemetry.context.ContextKey; +import io.opentelemetry.context.Scope; +import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.export.SimpleSpanProcessor; +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; + +class HandlerContextIntegrationTest { + private static final String TRACE_ID = "12345678901234567890123456789012"; + + @ParameterizedTest + @CsvSource({"true,success", "false,success", "true,failure", "false,failure", "true,suspension", "false,suspension" + }) + void reusedHandlerWorkerHasNoLeakedScope(boolean executionView, String outcome) throws Exception { + var exporter = InMemorySpanExporter.create(); + var builder = SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)); + var pluginConfig = OtelPluginConfig.builder() + .enableMdc(false) + .contextExtractor( + () -> new ExtractedContext(TRACE_ID, "1234567890123456", ExtractedContext.Sampling.SAMPLED)) + .build(); + DurableExecutionPlugin plugin = executionView + ? new ExecutionOtelPlugin(builder, pluginConfig) + : new InvocationOtelPlugin(builder, pluginConfig); + var executor = Executors.newSingleThreadExecutor(); + var config = DurableConfig.builder() + .withExecutorService(executor) + .withPlugins(plugin) + .build(); + try { + var runner = LocalDurableTestRunner.create( + String.class, + (input, ctx) -> { + assertEquals(TRACE_ID, Span.current().getSpanContext().getTraceId()); + if (outcome.equals("failure")) throw new IllegalStateException("user failure"); + if (outcome.equals("suspension")) ctx.wait("resume", Duration.ofSeconds(1)); + return "done"; + }, + config); + var result = runner.run("input"); + assertEquals( + switch (outcome) { + case "failure" -> ExecutionStatus.FAILED; + case "suspension" -> ExecutionStatus.PENDING; + default -> ExecutionStatus.SUCCEEDED; + }, + result.getStatus()); + assertFalse( + executor.submit(() -> Span.current().getSpanContext().isValid()) + .get(5, TimeUnit.SECONDS), + "invocation-end cleanup must restore the handler worker before the response returns"); + } finally { + executor.shutdownNow(); + } + } + + @ParameterizedTest + @CsvSource({"true,true", "false,true", "true,false", "false,false"}) + void preservesCompatibleAmbientAndRestoresUnrelatedAmbientAfterFailure(boolean executionView, boolean sameTrace) { + var exporter = InMemorySpanExporter.create(); + var builder = SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)); + var pluginConfig = OtelPluginConfig.builder() + .enableMdc(false) + .contextExtractor( + () -> new ExtractedContext(TRACE_ID, "1234567890123456", ExtractedContext.Sampling.SAMPLED)) + .build(); + DurableExecutionPlugin plugin = executionView + ? new ExecutionOtelPlugin(builder, pluginConfig) + : new InvocationOtelPlugin(builder, pluginConfig); + var ambient = SpanContext.create( + sameTrace ? TRACE_ID : "abcdefabcdefabcdefabcdefabcdefab", + "abcdefabcdefabcd", + TraceFlags.getSampled(), + TraceState.getDefault()); + var previous = Context.current(); + try (var ignored = Span.wrap(ambient).makeCurrent()) { + plugin.onInvocationStart(new InvocationInfo( + "req", + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/name/id", + true, + Instant.now())); + try { + var active = Span.current().getSpanContext(); + assertEquals(TRACE_ID, active.getTraceId()); + if (sameTrace) assertEquals(ambient, active); + else assertNotEquals(ambient.getSpanId(), active.getSpanId()); + } finally { + plugin.onInvocationEnd(new InvocationEndInfo( + "req", "arn", true, InvocationStatus.FAILED, new IllegalStateException("handler failure"))); + } + assertEquals(ambient, Span.current().getSpanContext()); + } + assertSame(previous, Context.current()); + } + + @ParameterizedTest + @CsvSource({"true,true", "false,true", "true,false", "false,false"}) + void rootContextIsValidAndRestoredAcrossNestedWorkAndResume(boolean executionView, boolean sampled) { + var exporter = InMemorySpanExporter.create(); + var builder = SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)); + var pluginConfig = OtelPluginConfig.builder() + .enableMdc(false) + .contextExtractor(() -> new ExtractedContext( + TRACE_ID, + "1234567890123456", + sampled ? ExtractedContext.Sampling.SAMPLED : ExtractedContext.Sampling.NOT_SAMPLED)) + .build(); + DurableExecutionPlugin plugin = executionView + ? new ExecutionOtelPlugin(builder, pluginConfig) + : new InvocationOtelPlugin(builder, pluginConfig); + var executor = Executors.newCachedThreadPool(); + var config = DurableConfig.builder() + .withExecutorService(executor) + .withPlugins(plugin) + .build(); + var userProvider = SdkTracerProvider.builder() + .addSpanProcessor(SimpleSpanProcessor.create(exporter)) + .build(); + var userTracer = userProvider.get("user"); + var roots = new ArrayList(); + var bodyCalls = new AtomicInteger(); + try { + var runner = LocalDurableTestRunner.create( + String.class, + (input, ctx) -> { + var root = Span.current().getSpanContext(); + assertTrue(root.isValid(), "root handler must have a valid context without an agent"); + assertEquals(TRACE_ID, root.getTraceId()); + assertEquals(sampled, root.isSampled()); + roots.add(root); + userTracer.spanBuilder("user-handler").startSpan().end(); + ctx.step("before", String.class, step -> { + bodyCalls.incrementAndGet(); + assertNotEquals( + root.getSpanId(), + Span.current().getSpanContext().getSpanId()); + return "before"; + }); + assertEquals(root, Span.current().getSpanContext()); + ctx.runInChildContext("child", String.class, child -> { + assertNotEquals( + root.getSpanId(), + Span.current().getSpanContext().getSpanId()); + return "child"; + }); + assertEquals(root, Span.current().getSpanContext()); + userTracer + .spanBuilder("user-handler-restored") + .startSpan() + .end(); + ctx.wait("resume", Duration.ofSeconds(1)); + assertEquals(root, Span.current().getSpanContext()); + userTracer + .spanBuilder("user-handler-after-resume") + .startSpan() + .end(); + return "done"; + }, + config); + var first = runner.run("input"); + assertEquals(ExecutionStatus.PENDING, first.getStatus()); + runner.advanceTime(); + var last = runner.run("input"); + assertEquals(ExecutionStatus.SUCCEEDED, last.getStatus()); + assertEquals(2, roots.size()); + assertEquals(1, bodyCalls.get()); + assertFalse(Span.current().getSpanContext().isValid()); + if (sampled) { + var rootName = executionView ? "Workflow" : "Invocation"; + var spans = exporter.getFinishedSpanItems(); + var userSpans = spans.stream() + .filter(span -> span.getName().startsWith("user-handler")) + .toList(); + assertEquals(5, userSpans.size()); + for (var span : userSpans) { + assertEquals(TRACE_ID, span.getTraceId()); + assertTrue(roots.stream().anyMatch(root -> root.getSpanId().equals(span.getParentSpanId()))); + } + for (var root : roots) { + assertTrue(spans.stream() + .anyMatch(s -> rootName.equals(s.getName()) + && root.getSpanId().equals(s.getSpanId()))); + } + } else { + assertTrue(exporter.getFinishedSpanItems().isEmpty()); + } + } finally { + executor.shutdownNow(); + userProvider.close(); + } + } + + @ParameterizedTest + @CsvSource({"true,true", "true,false", "false,true", "false,false"}) + void customPluginScopesUnwindBeforeOtelScopeOnTheReusedWorker(boolean executionView, boolean suspend) + throws Exception { + var key = ContextKey.named("application-context"); + var exporter = InMemorySpanExporter.create(); + var builder = SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)); + var settings = OtelPluginConfig.builder() + .enableMdc(false) + .contextExtractor( + () -> new ExtractedContext(TRACE_ID, "1234567890123456", ExtractedContext.Sampling.SAMPLED)) + .build(); + DurableExecutionPlugin delegate = executionView + ? new ExecutionOtelPlugin(builder, settings) + : new InvocationOtelPlugin(builder, settings); + var ends = new ArrayList(); + var otel = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + delegate.onInvocationStart(info); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add("otel"); + delegate.onInvocationEnd(info); + } + }; + var custom = new DurableExecutionPlugin() { + private Scope scope; + + @Override + public void onInvocationStart(InvocationInfo info) { + scope = Context.current().with(key, "active").makeCurrent(); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add("custom"); + assertEquals("active", Context.current().get(key)); + scope.close(); + } + }; + var workers = Executors.newSingleThreadExecutor(); + try { + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + assertEquals("active", Context.current().get(key)); + assertEquals(TRACE_ID, Span.current().getSpanContext().getTraceId()); + if (suspend) context.wait("pause", Duration.ofSeconds(1)); + return "done"; + }, + DurableConfig.builder() + .withExecutorService(workers) + .withPlugins(otel, custom) + .build()); + for (var invocation = 0; invocation < 2; invocation++) { + runner.run("input"); + assertEquals(List.of("custom", "otel"), ends); + ends.clear(); + assertFalse(workers.submit(() -> Span.current().getSpanContext().isValid()) + .get(3, TimeUnit.SECONDS)); + assertNull(workers.submit(() -> Context.current().get(key)).get(3, TimeUnit.SECONDS)); + runner.advanceTime(); + } + } finally { + workers.shutdownNow(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/HandlerMdcIntegrationTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/HandlerMdcIntegrationTest.java new file mode 100644 index 000000000..7fe9f18cf --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/HandlerMdcIntegrationTest.java @@ -0,0 +1,283 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import java.time.Duration; +import java.util.Arrays; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.concurrent.locks.LockSupport; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.ValueSource; +import org.slf4j.MDC; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; + +class HandlerMdcIntegrationTest { + private static final String TRACE_ID = "12345678901234567890123456789012"; + + @ParameterizedTest + @ValueSource(strings = {"success", "failure", "suspension", "inputFailure"}) + void noPluginPathRetainsExistingMdcBehavior(String outcome) throws Exception { + var executor = Executors.newSingleThreadExecutor(); + var callerBefore = MDC.getCopyOfContextMap(); + var callerMdc = Map.of(MdcSpanEnricher.MDC_TRACE_ID, "caller-trace", "caller", "retained"); + var workerMdc = ambientMdc(); + try { + executor.submit(() -> MDC.setContextMap(workerMdc)).get(5, TimeUnit.SECONDS); + MDC.setContextMap(callerMdc); + runWithoutPlugins(executor, outcome); + // The existing logger clears MDC once entered; failed input never enters that scope. + var expectedWorker = outcome.equals("inputFailure") ? workerMdc : Map.of(); + assertRestored(executor, expectedWorker, callerMdc); + } finally { + executor.shutdownNow(); + if (callerBefore == null) MDC.clear(); + else MDC.setContextMap(callerBefore); + } + } + + private static void runWithoutPlugins(ExecutorService executor, String outcome) { + var config = DurableConfig.builder() + .withExecutorService(executor) + .withSerDes(outcome.equals("inputFailure") ? failingInputSerDes() : new JacksonSerDes()) + .build(); + assertTrue(config.getPluginRunner().isEmpty()); + var runner = LocalDurableTestRunner.create( + String.class, + (input, ctx) -> { + if (outcome.equals("failure")) throw new IllegalStateException("handler failure"); + if (outcome.equals("suspension")) ctx.wait("pause", Duration.ofSeconds(1)); + return "done"; + }, + config); + var expected = + switch (outcome) { + case "failure", "inputFailure" -> ExecutionStatus.FAILED; + case "suspension" -> ExecutionStatus.PENDING; + default -> ExecutionStatus.SUCCEEDED; + }; + assertEquals(expected, runner.run("input").getStatus()); + } + + @ParameterizedTest + @CsvSource({ + "true,success,false", "false,success,false", "true,failure,false", "false,failure,false", + "true,suspension,false", "false,suspension,false", "true,inputFailure,false", "false,inputFailure,false", + "true,success,true", "false,success,true", "true,failure,true", "false,failure,true", + "true,suspension,true", "false,suspension,true", "true,inputFailure,true", "false,inputFailure,true" + }) + void preservesCallerAndReusedWorkerMdc(boolean executionView, String outcome, boolean ambient) throws Exception { + var executor = Executors.newSingleThreadExecutor(); + var callerBefore = MDC.getCopyOfContextMap(); + var callerMdc = Map.of(MdcSpanEnricher.MDC_TRACE_ID, "caller-trace", "caller", "retained"); + var workerMdc = ambient ? ambientMdc() : Map.of(); + try { + executor.submit(() -> MDC.setContextMap(workerMdc)).get(5, TimeUnit.SECONDS); + MDC.setContextMap(callerMdc); + runInvocations(executor, executionView, outcome, workerMdc, callerMdc); + } finally { + executor.shutdownNow(); + if (callerBefore == null) MDC.clear(); + else MDC.setContextMap(callerBefore); + } + } + + private static Map ambientMdc() { + return Map.of( + MdcSpanEnricher.MDC_TRACE_ID, + "worker-trace", + MdcSpanEnricher.MDC_SPAN_ID, + "worker-span", + MdcSpanEnricher.MDC_TRACE_SAMPLED, + "worker-sampled", + "application", + "retained"); + } + + private static void runInvocations( + ExecutorService executor, + boolean executionView, + String outcome, + Map workerMdc, + Map callerMdc) + throws Exception { + var calls = new AtomicInteger(); + var runner = LocalDurableTestRunner.create( + String.class, + (input, ctx) -> { + calls.incrementAndGet(); + assertEquals(TRACE_ID, MDC.get(MdcSpanEnricher.MDC_TRACE_ID)); + if (outcome.equals("failure")) throw new IllegalStateException("handler failure"); + if (outcome.equals("suspension")) ctx.wait("resume", Duration.ofSeconds(1)); + return "done"; + }, + config(executor, executionView, outcome)); + var first = runner.run("input"); + assertEquals( + switch (outcome) { + case "failure", "inputFailure" -> ExecutionStatus.FAILED; + case "suspension" -> ExecutionStatus.PENDING; + default -> ExecutionStatus.SUCCEEDED; + }, + first.getStatus()); + assertEquals(outcome.equals("inputFailure") ? 0 : 1, calls.get()); + assertRestored(executor, workerMdc, callerMdc); + if (outcome.equals("suspension")) { + runner.advanceTime(); + assertEquals(ExecutionStatus.SUCCEEDED, runner.run("input").getStatus()); + assertEquals(2, calls.get()); + assertRestored(executor, workerMdc, callerMdc); + } + } + + private static DurableConfig config(ExecutorService executor, boolean executionView, String outcome) { + var settings = OtelPluginConfig.builder() + .enableMdc(true) + .contextExtractor( + () -> new ExtractedContext(TRACE_ID, "1234567890123456", ExtractedContext.Sampling.SAMPLED)) + .build(); + DurableExecutionPlugin plugin = executionView + ? new ExecutionOtelPlugin(SdkTracerProvider.builder(), settings) + : new InvocationOtelPlugin(SdkTracerProvider.builder(), settings); + return DurableConfig.builder() + .withExecutorService(executor) + .withPlugins(plugin) + .withSerDes(outcome.equals("inputFailure") ? failingInputSerDes() : new JacksonSerDes()) + .build(); + } + + private static SerDes failingInputSerDes() { + return new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + @Override + public String serialize(Object value) { + return delegate.serialize(value); + } + + @Override + public T deserialize(String data, TypeToken type) { + if ("\"input\"".equals(data)) throw new IllegalStateException("input failure"); + return delegate.deserialize(data, type); + } + }; + } + + @ParameterizedTest + @ValueSource(strings = {"success", "failure", "inputFailure"}) + void endHookObservesTaskMdcBeforeReusedWorkerAmbientStateIsRestored(String outcome) throws Exception { + var worker = Executors.newSingleThreadExecutor(); + var callerThread = new AtomicReference(); + var caller = Executors.newSingleThreadExecutor(task -> { + var thread = new Thread(task, "mdc-end-caller"); + callerThread.set(thread); + return thread; + }); + var started = new CountDownLatch(1); + var release = new CountDownLatch(1); + var startThread = new AtomicReference(); + var endThread = new AtomicReference(); + var endMdc = new AtomicReference>(); + var ambient = Map.of("worker", "ambient"); + var plugin = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + startThread.set(Thread.currentThread()); + MDC.put("start-hook", "request"); + started.countDown(); + awaitMdcLatch(release); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + endThread.set(Thread.currentThread()); + var values = MDC.getCopyOfContextMap(); + endMdc.set(values == null ? Map.of() : values); + MDC.put("end-hook", "must-not-leak"); + } + }; + try { + worker.submit(() -> MDC.setContextMap(ambient)).get(3, TimeUnit.SECONDS); + var config = DurableConfig.builder() + .withExecutorService(worker) + .withPlugins(plugin) + .withSerDes(outcome.equals("inputFailure") ? failingInputSerDes() : new JacksonSerDes()) + .build(); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + if (outcome.equals("failure")) throw new IllegalStateException("body failure"); + return "done"; + }, + config); + var result = caller.submit(() -> runner.run("input")); + assertTrue(started.await(3, TimeUnit.SECONDS)); + awaitMdcCallerJoin(callerThread.get()); + release.countDown(); + assertEquals( + outcome.equals("success") ? ExecutionStatus.SUCCEEDED : ExecutionStatus.FAILED, + result.get(5, TimeUnit.SECONDS).getStatus()); + assertSame(startThread.get(), endThread.get()); + // Input failure skips DurableLogger; entered handlers retain its established MDC-clearing behavior. + assertEquals( + outcome.equals("inputFailure") ? Map.of("worker", "ambient", "start-hook", "request") : Map.of(), + endMdc.get(), + "End hooks must not see prematurely restored ambient worker MDC"); + assertEquals(ambient, worker.submit(MDC::getCopyOfContextMap).get(3, TimeUnit.SECONDS)); + } finally { + release.countDown(); + caller.shutdownNow(); + worker.shutdownNow(); + assertTrue(caller.awaitTermination(3, TimeUnit.SECONDS)); + assertTrue(worker.awaitTermination(3, TimeUnit.SECONDS)); + } + } + + private static void awaitMdcLatch(CountDownLatch latch) { + try { + assertTrue(latch.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throw new AssertionError(error); + } + } + + private static void awaitMdcCallerJoin(Thread caller) { + var deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(3); + while (System.nanoTime() < deadline) { + if (caller.getState() == Thread.State.WAITING + && Arrays.stream(caller.getStackTrace()) + .anyMatch(frame -> frame.getClassName().equals(CompletableFuture.class.getName()) + && frame.getMethodName().equals("join"))) return; + LockSupport.parkNanos(TimeUnit.MILLISECONDS.toNanos(1)); + } + fail("Caller did not attach its completion observer before releasing the handler"); + } + + private static void assertRestored( + ExecutorService executor, Map workerMdc, Map callerMdc) throws Exception { + var callerAfter = MDC.getCopyOfContextMap(); + var after = executor.submit(MDC::getCopyOfContextMap).get(5, TimeUnit.SECONDS); + assertAll( + () -> assertEquals(callerMdc, callerAfter, "finalization must preserve caller MDC"), + () -> assertEquals(workerMdc, after == null ? Map.of() : after, "worker must regain its original MDC")); + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationClockContainmentTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationClockContainmentTest.java new file mode 100644 index 000000000..da378fd86 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationClockContainmentTest.java @@ -0,0 +1,145 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import io.opentelemetry.sdk.common.Clock; +import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.export.SimpleSpanProcessor; +import java.time.Instant; +import java.util.concurrent.atomic.AtomicLong; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.EnumSource; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.plugin.OperationEndInfo; +import software.amazon.lambda.durable.plugin.OperationInfo; + +class InvocationClockContainmentTest { + @ParameterizedTest + @EnumSource(InvocationStatus.class) + void openChildAndParentShareTheConfiguredSdkClock(InvocationStatus status) { + var wall = new AtomicLong(100_000_000_000L); + var monotonic = new AtomicLong(); + var clock = new Clock() { + public long now() { + return wall.addAndGet(1_000_000); + } + + public long nanoTime() { + return monotonic.addAndGet(1_000); + } + }; + try (var exporter = InMemorySpanExporter.create()) { + var plugin = new InvocationOtelPlugin( + SdkTracerProvider.builder().setClock(clock).addSpanProcessor(SimpleSpanProcessor.create(exporter)), + OtelPluginConfig.builder() + .contextExtractor(() -> null) + .enableMdc(false) + .build()); + var start = Instant.ofEpochSecond(10); + var arn = "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/clock/execution"; + plugin.onInvocationStart(new InvocationInfo("request", arn, true, start)); + plugin.onOperationStart(new OperationInfo( + "parent", "parent", "CONTEXT", "WaitForCallback", null, start, null, null, false)); + plugin.onOperationStart( + new OperationInfo("child", "child", "CALLBACK", "Callback", "parent", start, null, null, false)); + plugin.onInvocationEnd(new InvocationEndInfo( + "request", + arn, + true, + status, + status == InvocationStatus.FAILED ? new IllegalStateException("failure") : null)); + var spans = exporter.getFinishedSpanItems(); + var child = spans.stream() + .filter(s -> s.getName().equals("child")) + .findFirst() + .orElseThrow(); + var parent = spans.stream() + .filter(s -> s.getName().equals("parent")) + .findFirst() + .orElseThrow(); + var invocation = spans.stream() + .filter(s -> s.getName().equals("Invocation")) + .findFirst() + .orElseThrow(); + assertEquals(parent.getSpanId(), child.getParentSpanId()); + assertEquals( + parent.getSpanContext(), + child.getParentSpanContext(), + "Live parent keeps the cached trace ID, span ID, flags and trace state"); + assertEquals(invocation.getSpanId(), parent.getParentSpanId()); + assertTrue(spans.indexOf(child) < spans.indexOf(parent), "The plugin already drains children first"); + assertTrue( + child.getEndEpochNanos() <= parent.getEndEpochNanos(), + "Separate anchor clocks must not invert child-first end timestamps"); + assertTrue(parent.getEndEpochNanos() <= invocation.getEndEpochNanos()); + assertTrue(child.getStartEpochNanos() >= parent.getStartEpochNanos()); + assertTrue( + invocation.getEndEpochNanos() < 101_000_000_000L, + "Keep the configured clock; do not substitute system time or relax containment"); + } + } + + @Test + void endedParentRetainsItsCachedContextForAReplayedChildWithoutReexport() { + try (var exporter = InMemorySpanExporter.create()) { + var plugin = new InvocationOtelPlugin( + SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)), + OtelPluginConfig.builder() + .contextExtractor(() -> null) + .enableMdc(false) + .build()); + var start = Instant.now(); + var arn = "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/fallback/execution"; + plugin.onInvocationStart(new InvocationInfo("request", arn, false, start)); + plugin.onOperationStart(new OperationInfo( + "parent", "parent", "CONTEXT", "RunInChildContext", null, start, null, null, true)); + plugin.onOperationEnd(new OperationEndInfo( + "parent", + "parent", + "CONTEXT", + "RunInChildContext", + null, + start, + Instant.now(), + "SUCCEEDED", + null, + true, + null, + null)); + var parent = exporter.getFinishedSpanItems().stream() + .filter(s -> s.getName().equals("parent")) + .findFirst() + .orElseThrow(); + plugin.onOperationEnd(new OperationEndInfo( + "child", + "child", + "CALLBACK", + "Callback", + "parent", + start, + Instant.now(), + "SUCCEEDED", + null, + true, + null, + null)); + plugin.onInvocationEnd(new InvocationEndInfo("request", arn, false, InvocationStatus.PENDING, null)); + var spans = exporter.getFinishedSpanItems(); + var child = spans.stream() + .filter(s -> s.getName().equals("child")) + .findFirst() + .orElseThrow(); + assertEquals(parent.getSpanContext(), child.getParentSpanContext()); + assertEquals( + 1, spans.stream().filter(s -> s.getName().equals("parent")).count()); + assertEquals( + 1, spans.stream().filter(s -> s.getName().equals("child")).count()); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndCompatibilityTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndCompatibilityTest.java new file mode 100644 index 000000000..9cb19b13f --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndCompatibilityTest.java @@ -0,0 +1,303 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.*; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.concurrent.locks.LockSupport; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +/** Verifies invocation-hook thread affinity, including executors that complete tasks before submission returns. */ +class InvocationEndCompatibilityTest { + @ParameterizedTest + @CsvSource({ + "false,false,false", "true,false,false", "false,true,false", "true,true,false", + "false,false,true", "true,false,true", "false,true,true", "true,true,true" + }) + void invocationEndRunsOnHandlerThread(boolean mixed, boolean precompleted, boolean pending) throws Exception { + var local = new ThreadLocal(); + var observations = new ArrayList(); + var callerThread = new AtomicReference(); + var caller = Executors.newSingleThreadExecutor(task -> { + var thread = daemon(task, "legacy-probe-caller"); + callerThread.set(thread); + return thread; + }); + var workers = + new ThreadPoolExecutor( + 1, + 1, + 0, + TimeUnit.SECONDS, + new LinkedBlockingQueue<>(), + task -> daemon(task, "legacy-probe-worker")) { + @Override + public void execute(Runnable task) { + var finished = new CountDownLatch(1); + super.execute(() -> { + try { + task.run(); + } finally { + finished.countDown(); + } + }); + if (precompleted) await(finished); + } + }; + try { + for (int round = 1; round <= 2; round++) { + var marker = "inv" + round; + var entered = new CountDownLatch(1); + var release = new CountDownLatch(1); + var endValue = new AtomicReference(); + var startThread = new AtomicReference(); + var endThread = new AtomicReference(); + var previous = new AtomicReference(); + var endCalls = new AtomicInteger(); + var endOrder = new CopyOnWriteArrayList(); + var legacy = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + startThread.set(Thread.currentThread()); + previous.set(local.get()); + local.set(marker); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + endCalls.incrementAndGet(); + endOrder.add("legacy"); + endThread.set(Thread.currentThread()); + endValue.set(local.get()); + local.remove(); + } + }; + var plugins = new ArrayList(); + plugins.add(legacy); + if (mixed) + plugins.add( + 0, + new InvocationOtelPlugin( + SdkTracerProvider.builder(), + OtelPluginConfig.builder().enableMdc(false).build()) { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + endOrder.add("otel"); + super.onInvocationEnd(info); + } + }); + var config = DurableConfig.builder() + .withExecutorService(workers) + .withDurableExecutionClient(new LocalMemoryExecutionClient()) + .withCheckpointDelay(Duration.ZERO) + .withPlugins(plugins.toArray(DurableExecutionPlugin[]::new)) + .build(); + var result = caller.submit(() -> DurableExecutor.execute( + input(marker), + null, + TypeToken.get(String.class), + (value, context) -> { + entered.countDown(); + if (!precompleted) await(release); + if (pending) context.wait("pause", Duration.ofSeconds(1)); + return "done"; + }, + config)); + try { + assertTrue(entered.await(3, TimeUnit.SECONDS)); + if (!precompleted) { + awaitCallerJoin(callerThread.get()); + release.countDown(); + } + assertEquals( + pending ? ExecutionStatus.PENDING : ExecutionStatus.SUCCEEDED, + result.get(5, TimeUnit.SECONDS).status()); + var after = workers.submit(local::get).get(3, TimeUnit.SECONDS); + observations.add(new Observation( + startThread.get(), endThread.get(), endValue.get(), after, previous.get(), endCalls.get())); + assertEquals( + mixed ? List.of("legacy", "otel") : List.of("legacy"), + endOrder, + "all end hooks unwind registration order and execute once"); + } finally { + release.countDown(); + } + } + for (int index = 0; index < observations.size(); index++) { + var observation = observations.get(index); + assertEquals(1, observation.endCalls()); + assertSame(observation.startThread(), observation.endThread()); + assertNotSame(callerThread.get(), observation.endThread()); + assertEquals("inv" + (index + 1), observation.endValue()); + assertNull(observation.workerAfter(), "end hook clears its original worker value"); + assertNull(observation.previous(), "reused worker must not retain the prior invocation"); + } + } finally { + caller.shutdownNow(); + workers.shutdownNow(); + assertTrue(caller.awaitTermination(3, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + } + } + + @Test + void suspendedInvocationWaitsForEndOnTheReusedWorker() throws Exception { + var local = new ThreadLocal(); + var callerThread = new AtomicReference(); + var callers = Executors.newSingleThreadExecutor(task -> { + var thread = daemon(task, "scoped-handoff-caller"); + callerThread.set(thread); + return thread; + }); + var workers = Executors.newSingleThreadExecutor(task -> daemon(task, "scoped-handoff-worker")); + try { + for (var round = 0; round < 3; round++) { + var marker = "scoped-" + round; + var scoped = new PausingOtelPlugin(); + var previous = new AtomicReference(); + var observed = new AtomicReference(); + var owner = new AtomicReference(); + var endThread = new AtomicReference(); + var legacy = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + owner.set(Thread.currentThread()); + previous.set(local.get()); + local.set(marker); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + endThread.set(Thread.currentThread()); + observed.set(local.get()); + local.remove(); + } + }; + var config = DurableConfig.builder() + .withExecutorService(workers) + .withDurableExecutionClient(new LocalMemoryExecutionClient()) + .withCheckpointDelay(Duration.ZERO) + .withPlugins(scoped, legacy) + .build(); + var result = callers.submit(() -> DurableExecutor.execute( + input(marker), + null, + TypeToken.get(String.class), + (value, context) -> { + context.wait("pause", Duration.ofSeconds(1)); + return "done"; + }, + config)); + try { + assertTrue(scoped.closeEntered.await(3, TimeUnit.SECONDS)); + assertThrows( + TimeoutException.class, + () -> result.get(600, TimeUnit.MILLISECONDS), + "a blocked end hook has no fallback to the caller thread"); + scoped.releaseClose.countDown(); + assertEquals( + ExecutionStatus.PENDING, + result.get(5, TimeUnit.SECONDS).status()); + assertSame(owner.get(), endThread.get()); + assertEquals(marker, observed.get()); + assertNull(previous.get(), "The prior invocation's end hook must clear the reused worker"); + assertNull(workers.submit(local::get).get(3, TimeUnit.SECONDS)); + } finally { + scoped.releaseClose.countDown(); + } + } + } finally { + callers.shutdownNow(); + workers.shutdownNow(); + assertTrue(callers.awaitTermination(3, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + } + } + + private static final class PausingOtelPlugin extends InvocationOtelPlugin { + private final CountDownLatch closeEntered = new CountDownLatch(1); + private final CountDownLatch releaseClose = new CountDownLatch(1); + + private PausingOtelPlugin() { + super( + SdkTracerProvider.builder(), + OtelPluginConfig.builder().enableMdc(false).build()); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + try { + super.onInvocationEnd(info); + } finally { + closeEntered.countDown(); + await(releaseClose); + } + } + } + + private record Observation( + Thread startThread, Thread endThread, String endValue, String workerAfter, String previous, int endCalls) {} + + private static void awaitCallerJoin(Thread caller) { + var end = System.nanoTime() + TimeUnit.SECONDS.toNanos(3); + while (System.nanoTime() < end) { + if (caller.getState() == Thread.State.WAITING + && Arrays.stream(caller.getStackTrace()) + .anyMatch(frame -> frame.getClassName().equals(CompletableFuture.class.getName()) + && frame.getMethodName().equals("join"))) return; + LockSupport.parkNanos(TimeUnit.MILLISECONDS.toNanos(1)); + } + fail("invocation caller did not reach its future join before releasing the handler"); + } + + private static DurableExecutionInput input(String marker) { + var operation = Operation.builder() + .id("id") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/" + marker + "/id", + "token", + CheckpointUpdatedExecutionState.builder().operations(operation).build()); + } + + private static Thread daemon(Runnable task, String name) { + var thread = new Thread(task, name); + thread.setDaemon(true); + return thread; + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throw new AssertionError(error); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndFailureTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndFailureTest.java new file mode 100644 index 000000000..cff654b99 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndFailureTest.java @@ -0,0 +1,218 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import io.opentelemetry.api.trace.Span; +import io.opentelemetry.api.trace.SpanContext; +import io.opentelemetry.api.trace.TraceFlags; +import io.opentelemetry.api.trace.TraceState; +import io.opentelemetry.context.Context; +import io.opentelemetry.context.Scope; +import io.opentelemetry.sdk.common.CompletableResultCode; +import io.opentelemetry.sdk.trace.ReadWriteSpan; +import io.opentelemetry.sdk.trace.ReadableSpan; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.SpanProcessor; +import java.time.Instant; +import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.plugin.PluginRunner; + +class InvocationEndFailureTest { + @ParameterizedTest + @MethodSource("failures") + @SuppressWarnings("removal") + void restoresContextWhenInvocationEndFails(boolean executionView, String phase, String failureKind) + throws Exception { + Throwable failure = + switch (failureKind) { + case "exception" -> new IllegalStateException("telemetry failed"); + case "linkage" -> new NoClassDefFoundError("incompatible telemetry dependency"); + case "fatal" -> new InternalError("fatal telemetry failure"); + default -> new ThreadDeath(); + }; + var processor = new FailingProcessor(phase, failure); + var builder = SdkTracerProvider.builder().addSpanProcessor(processor); + var config = OtelPluginConfig.builder() + .enableMdc(false) + .contextExtractor(() -> new ExtractedContext( + "12345678901234567890123456789012", "1234567890123456", ExtractedContext.Sampling.SAMPLED)) + .build(); + DurableExecutionPlugin plugin = + executionView ? new ExecutionOtelPlugin(builder, config) : new InvocationOtelPlugin(builder, config); + var healthyEnds = new AtomicInteger(); + var healthy = new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + healthyEnds.incrementAndGet(); + } + }; + var runner = new PluginRunner(List.of(healthy, plugin)); + var ambient = Span.wrap(SpanContext.create( + "abcdefabcdefabcdefabcdefabcdefab", + "abcdefabcdefabcd", + TraceFlags.getSampled(), + TraceState.getDefault())); + try (var ignored = ambient.makeCurrent()) { + var original = Context.current(); + runner.onInvocationStart(new InvocationInfo("req", "arn", true, Instant.now())); + assertNotSame(original, Context.current()); + if (phase.equals("scope")) { + var field = plugin.getClass().getDeclaredField("handlerScope"); + field.setAccessible(true); + var scope = (Scope) field.get(plugin); + field.set(plugin, (Scope) () -> { + scope.close(); + raise(failure); + }); + } + var end = new InvocationEndInfo("req", "arn", true, InvocationStatus.SUCCEEDED, null); + var fatal = failureKind.equals("fatal") || failureKind.equals("thread-death"); + if (fatal) assertSame(failure, assertThrows(Error.class, () -> runner.onInvocationEnd(end))); + else assertDoesNotThrow(() -> runner.onInvocationEnd(end)); + assertSame(original, Context.current(), "telemetry errors must not leave the handler context attached"); + assertEquals(1, healthyEnds.get(), "remaining invocation-end hooks must always release their resources"); + // Cleanup state is consumed even if the scope's close implementation throws. + assertDoesNotThrow(() -> plugin.onInvocationEnd(end)); + assertSame(original, Context.current()); + } + } + + @ParameterizedTest + @MethodSource("combinedFailures") + void preservesCombinedFinalizationAndScopeFailures(boolean executionView, String phase, String pair) + throws Exception { + Throwable primary = + switch (pair.split("-")[0]) { + case "fatal" -> new InternalError("primary finalization"); + case "assert" -> new AssertionError("primary finalization"); + case "linkage" -> new NoClassDefFoundError("primary finalization"); + default -> new IllegalStateException("primary finalization"); + }; + Throwable cleanup = + switch (pair.split("-")[1]) { + case "fatal" -> new InternalError("scope cleanup"); + case "assert" -> new AssertionError("scope cleanup"); + case "same" -> primary; + default -> new IllegalStateException("scope cleanup"); + }; + var cleanupWins = List.of("runtime-fatal", "assert-fatal", "runtime-assert", "linkage-assert") + .contains(pair); + var expected = cleanupWins ? cleanup : primary; + var secondary = cleanupWins ? primary : cleanup; + var processor = new FailingProcessor(phase, primary); + var config = OtelPluginConfig.builder() + .enableMdc(false) + .contextExtractor(() -> new ExtractedContext( + "12345678901234567890123456789012", "1234567890123456", ExtractedContext.Sampling.SAMPLED)) + .build(); + var builder = SdkTracerProvider.builder().addSpanProcessor(processor); + DurableExecutionPlugin plugin = + executionView ? new ExecutionOtelPlugin(builder, config) : new InvocationOtelPlugin(builder, config); + var healthyEnds = new AtomicInteger(); + var closed = new AtomicInteger(); + var healthy = new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + healthyEnds.incrementAndGet(); + } + }; + var runner = new PluginRunner(List.of(healthy, plugin)); + var ambient = Span.wrap(SpanContext.create( + "abcdefabcdefabcdefabcdefabcdefab", + "abcdefabcdefabcd", + TraceFlags.getSampled(), + TraceState.getDefault())); + try (var ignored = ambient.makeCurrent()) { + var original = Context.current(); + runner.onInvocationStart(new InvocationInfo("req", "arn", true, Instant.now())); + var field = plugin.getClass().getDeclaredField("handlerScope"); + field.setAccessible(true); + var scope = (Scope) field.get(plugin); + field.set(plugin, (Scope) () -> { + closed.incrementAndGet(); + scope.close(); + raise(cleanup); + }); + var end = new InvocationEndInfo("req", "arn", true, InvocationStatus.SUCCEEDED, null); + if (expected instanceof RuntimeException || expected instanceof LinkageError) { + assertDoesNotThrow(() -> runner.onInvocationEnd(end)); + } else assertSame(expected, assertThrows(Error.class, () -> runner.onInvocationEnd(end))); + assertEquals(expected == secondary ? List.of() : List.of(secondary), List.of(expected.getSuppressed())); + assertSame(original, Context.current()); + assertEquals(1, healthyEnds.get()); + assertEquals(1, closed.get()); + assertDoesNotThrow(() -> plugin.onInvocationEnd(end)); + assertEquals(1, closed.get(), "scope cleanup is one-shot even when both phases throw"); + } + } + + private static Stream combinedFailures() { + return Stream.of(false, true) + .flatMap(view -> Stream.of("span", "flush") + .flatMap(phase -> Stream.of( + "fatal-runtime", + "assert-runtime", + "runtime-fatal", + "assert-fatal", + "fatal-fatal", + "fatal-same", + "runtime-assert", + "linkage-assert", + "runtime-runtime") + .map(pair -> Arguments.of(view, phase, pair)))); + } + + private static Stream failures() { + return Stream.of(false, true) + .flatMap(executionView -> Stream.of("span", "flush", "scope") + .flatMap(phase -> Stream.of("exception", "linkage", "fatal", "thread-death") + .map(kind -> Arguments.of(executionView, phase, kind)))); + } + + private static void raise(Throwable failure) { + if (failure instanceof RuntimeException exception) throw exception; + throw (Error) failure; + } + + private record FailingProcessor(String phase, Throwable failure) implements SpanProcessor { + @Override + public void onStart(Context parent, ReadWriteSpan span) {} + + @Override + public boolean isStartRequired() { + return false; + } + + @Override + public void onEnd(ReadableSpan span) { + if (phase.equals("span")) raise(failure); + } + + @Override + public boolean isEndRequired() { + return true; + } + + @Override + public CompletableResultCode forceFlush() { + if (phase.equals("flush")) raise(failure); + return CompletableResultCode.ofSuccess(); + } + + @Override + public CompletableResultCode shutdown() { + return CompletableResultCode.ofSuccess(); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndFinalizationTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndFinalizationTest.java new file mode 100644 index 000000000..5da60b868 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationEndFinalizationTest.java @@ -0,0 +1,195 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import io.opentelemetry.api.trace.Span; +import io.opentelemetry.context.Context; +import io.opentelemetry.sdk.common.CompletableResultCode; +import io.opentelemetry.sdk.trace.ReadWriteSpan; +import io.opentelemetry.sdk.trace.ReadableSpan; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.SpanProcessor; +import java.time.Duration; +import java.time.Instant; +import java.util.concurrent.*; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class InvocationEndFinalizationTest { + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void waitsForHandlerFinallyAndRealOtelFlushBeforeResponding(boolean executionView, boolean retry) throws Exception { + var flush = new BlockingFlush(); + var builder = SdkTracerProvider.builder().addSpanProcessor(flush); + var settings = OtelPluginConfig.builder() + .enableMdc(false) + .contextExtractor(() -> new ExtractedContext( + "12345678901234567890123456789012", "1234567890123456", ExtractedContext.Sampling.SAMPLED)) + .build(); + DurableExecutionPlugin delegate = executionView + ? new ExecutionOtelPlugin(builder, settings) + : new InvocationOtelPlugin(builder, settings); + var owner = new AtomicReference(); + var endOwner = new AtomicReference(); + var ends = new AtomicInteger(); + var plugin = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + owner.set(Thread.currentThread()); + delegate.onInvocationStart(info); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + endOwner.set(Thread.currentThread()); + delegate.onInvocationEnd(info); + assertFalse( + Span.current().getSpanContext().isValid(), "End must restore ambient context before returning"); + ends.incrementAndGet(); + } + }; + var workers = Executors.newCachedThreadPool(); + var callers = Executors.newSingleThreadExecutor(); + var enteredFinally = new CountDownLatch(1); + var releaseFinally = new CountDownLatch(1); + var original = new UnrecoverableDurableExecutionException( + ErrorObject.builder().errorMessage("original retry").build(), true); + var config = DurableConfig.builder() + .withDurableExecutionClient(new LocalMemoryExecutionClient()) + .withExecutorService(workers) + .withCheckpointDelay(Duration.ZERO) + .withPlugins(plugin) + .build(); + try { + var response = callers.submit(() -> DurableExecutor.execute( + input(), + null, + TypeToken.get(String.class), + (value, ctx) -> { + assertSame(owner.get(), Thread.currentThread()); + assertTrue(Span.current().getSpanContext().isValid()); + try { + if (retry) + ctx.step("retry", String.class, step -> { + throw original; + }); + else ctx.wait("pause", Duration.ofSeconds(1)); + return "done"; + } finally { + enteredFinally.countDown(); + await(releaseFinally); + } + }, + config)); + assertTrue(enteredFinally.await(3, TimeUnit.SECONDS)); + assertThrows( + TimeoutException.class, + () -> response.get(600, TimeUnit.MILLISECONDS), + "suspension or termination cannot bypass a blocked handler finally"); + assertEquals(0, flush.calls.get(), "End must wait for the handler to unwind"); + assertEquals(0, ends.get()); + releaseFinally.countDown(); + assertTrue(flush.entered.await(3, TimeUnit.SECONDS)); + assertSame(owner.get(), endOwner.get()); + assertSame(owner.get(), flush.owner.get()); + assertThrows( + TimeoutException.class, + () -> response.get(100, TimeUnit.MILLISECONDS), + "the invocation response must wait for the plugin's actual forceFlush result"); + flush.result.succeed(); + if (retry) { + var thrown = assertThrows(ExecutionException.class, () -> response.get(3, TimeUnit.SECONDS)); + assertSame(original, thrown.getCause()); + } else { + assertEquals( + ExecutionStatus.PENDING, + response.get(3, TimeUnit.SECONDS).status()); + } + assertEquals(1, flush.calls.get()); + assertEquals(1, ends.get()); + } finally { + releaseFinally.countDown(); + flush.result.succeed(); + workers.shutdown(); + callers.shutdown(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + assertTrue(callers.awaitTermination(3, TimeUnit.SECONDS)); + } + } + + private static DurableExecutionInput input() { + var operation = Operation.builder() + .id("id") + .name("test") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.now()) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/name/id", + "token", + CheckpointUpdatedExecutionState.builder().operations(operation).build()); + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(5, TimeUnit.SECONDS)); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throw new AssertionError(error); + } + } + + private static final class BlockingFlush implements SpanProcessor { + final AtomicInteger calls = new AtomicInteger(); + final AtomicReference owner = new AtomicReference<>(); + final CountDownLatch entered = new CountDownLatch(1); + final CompletableResultCode result = new CompletableResultCode(); + + @Override + public CompletableResultCode forceFlush() { + calls.incrementAndGet(); + owner.set(Thread.currentThread()); + entered.countDown(); + return result; + } + + @Override + public void onStart(Context parent, ReadWriteSpan span) {} + + @Override + public boolean isStartRequired() { + return false; + } + + @Override + public void onEnd(ReadableSpan span) {} + + @Override + public boolean isEndRequired() { + return false; + } + + @Override + public CompletableResultCode shutdown() { + return CompletableResultCode.ofSuccess(); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationOtelPluginTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationOtelPluginTest.java index d8d7e7f43..79926ac69 100644 --- a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationOtelPluginTest.java +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationOtelPluginTest.java @@ -308,8 +308,8 @@ void invocationStart_staysOnExecutionTrace_withoutLinkingAmbientSpan() { try (var ignored = Span.wrap(ambientSpanContext).makeCurrent()) { plugin.onInvocationStart(new InvocationInfo("req-1", "arn:exec1", true, Instant.now())); + plugin.onInvocationEnd(new InvocationEndInfo("req-1", "arn:exec1", true, InvocationStatus.SUCCEEDED, null)); } - plugin.onInvocationEnd(new InvocationEndInfo("req-1", "arn:exec1", true, InvocationStatus.SUCCEEDED, null)); var invocationSpan = spanByName("Invocation"); var workflowSpan = spanByName("Workflow"); diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationOutcomeBoundaryTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationOutcomeBoundaryTest.java new file mode 100644 index 000000000..f663951ab --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/InvocationOutcomeBoundaryTest.java @@ -0,0 +1,621 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; +import static software.amazon.lambda.durable.otel.SpanAttributes.DURABLE_INVOCATION_STATUS; + +import io.opentelemetry.api.trace.StatusCode; +import io.opentelemetry.context.Context; +import io.opentelemetry.context.ContextKey; +import io.opentelemetry.context.Scope; +import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.export.SimpleSpanProcessor; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Proxy; +import java.time.Duration; +import java.util.List; +import java.util.concurrent.CompletionException; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.Executors; +import java.util.concurrent.SynchronousQueue; +import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.ValueSource; +import org.slf4j.MDC; +import org.slf4j.spi.MDCAdapter; +import software.amazon.awssdk.services.lambda.model.ErrorObject; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; +import software.amazon.lambda.durable.util.ExceptionHelper; + +class InvocationOutcomeBoundaryTest { + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void cleanupWaitContractAlsoAppliesWithoutPlugins(boolean withPlugin, boolean retry) throws Exception { + var entered = new CountDownLatch(1); + var release = new CountDownLatch(1); + var original = retryError(); + var config = DurableConfig.builder(); + if (withPlugin) config.withPlugins(new DurableExecutionPlugin() {}); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + try { + if (retry) + context.step("retry", String.class, step -> { + throw original; + }); + else context.wait("pause", Duration.ofSeconds(1)); + return "done"; + } finally { + entered.countDown(); + await(release); + } + }, + config.build()); + var caller = Executors.newSingleThreadExecutor(); + try { + var response = caller.submit(() -> runner.run("input")); + assertTrue(entered.await(5, TimeUnit.SECONDS)); + assertThrows(TimeoutException.class, () -> response.get(700, TimeUnit.MILLISECONDS)); + release.countDown(); + if (retry) + assertSame( + original, + assertThrows(ExecutionException.class, () -> response.get(5, TimeUnit.SECONDS)) + .getCause()); + else + assertEquals( + ExecutionStatus.PENDING, + response.get(5, TimeUnit.SECONDS).getStatus()); + } finally { + release.countDown(); + caller.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void invocationEndFatalEscapesItsActualWorkerAfterSettlingCaller(boolean death, boolean earlierAssertion) + throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("fatal end cleanup"); + var ownerFailure = new AtomicReference(); + var escaped = new CountDownLatch(1); + var ends = new AtomicInteger(); + var assertion = new AssertionError("earlier nonfatal cleanup"); + var fatalPlugin = new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.incrementAndGet(); + throw fatal; + } + }; + var assertionPlugin = new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + throw assertion; + } + }; + var workers = Executors.newSingleThreadExecutor(task -> { + var worker = new Thread(task, "end-fatal-owner"); + worker.setUncaughtExceptionHandler((thread, failure) -> { + ownerFailure.set(failure); + escaped.countDown(); + }); + return worker; + }); + try { + var runner = LocalDurableTestRunner.create( + String.class, + (input, ctx) -> "done", + DurableConfig.builder() + .withExecutorService(workers) + .withPlugins( + earlierAssertion + ? new DurableExecutionPlugin[] {fatalPlugin, assertionPlugin} + : new DurableExecutionPlugin[] {fatalPlugin}) + .build()); + assertSame(fatal, assertThrows(Error.class, () -> runner.run("input"))); + assertTrue(escaped.await(2, TimeUnit.SECONDS), "fatal end cleanup must escape its actual worker"); + assertSame(fatal, ownerFailure.get()); + assertEquals(earlierAssertion ? List.of(assertion) : List.of(), List.of(fatal.getSuppressed())); + assertEquals(1, ends.get()); + } finally { + workers.shutdownNow(); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @ParameterizedTest + @CsvSource({"SUCCEEDED,false", "SUCCEEDED,true", "PENDING,false", "PENDING,true", "RETRYING,false", "RETRYING,true" + }) + void inlineNonfatalMdcRestorationPreservesSelectedOutcome(String outcome, boolean ambient) throws Exception { + runInlineMdcFailure(outcome, ambient, new IllegalStateException("restore failed"), null); + } + + @ParameterizedTest + @CsvSource({ + "SUCCEEDED,false,assertion", "SUCCEEDED,true,assertion", + "PENDING,false,assertion", "PENDING,true,assertion", + "RETRYING,false,assertion", "RETRYING,true,assertion", + "SUCCEEDED,false,linkage", "SUCCEEDED,true,linkage", + "PENDING,false,linkage", "PENDING,true,linkage", + "RETRYING,false,linkage", "RETRYING,true,linkage" + }) + void inlineNonfatalErrorRestorationPreservesSelectedOutcome(String outcome, boolean ambient, String kind) + throws Exception { + Error failure = kind.equals("linkage") + ? new NoClassDefFoundError("restore failed") + : new AssertionError("restore failed"); + runInlineMdcFailure(outcome, ambient, failure, null); + } + + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void inlineFatalRestorationStillEscapesWithOriginalIdentity(boolean death, boolean wrapped) throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("fatal restoration"); + runInlineMdcFailure( + "SUCCEEDED", false, wrapped ? new CompletionException(new ExecutionException(fatal)) : fatal, fatal); + } + + @ParameterizedTest + @ValueSource(strings = {"unreadable", "cycle", "fatal"}) + void lifecycleDiagnosticTraversalIsBoundedAndPreservesFatalIdentity(String kind) throws Exception { + var reads = new AtomicInteger(); + var fatal = new InternalError("fatal diagnostic"); + var failure = new CompletionException("diagnostic", null) { + @Override + public synchronized Throwable getCause() { + reads.incrementAndGet(); + if (kind.equals("cycle")) return this; + if (kind.equals("fatal")) throw fatal; + throw new IllegalStateException("unreadable diagnostic"); + } + }; + runInlineMdcFailure("SUCCEEDED", false, failure, kind.equals("fatal") ? fatal : null); + assertEquals(1, reads.get()); + } + + private static void runInlineMdcFailure( + String outcome, boolean ambient, Throwable restoreFailure, Error expectedFatal) throws Exception { + var adapter = MDC.getMDCAdapter(); + var before = MDC.getCopyOfContextMap(); + var ends = new CopyOnWriteArrayList(); + var injected = new AtomicBoolean(); + var original = retryError(); + var executor = new ThreadPoolExecutor(1, 1, 0, TimeUnit.SECONDS, new SynchronousQueue()) { + @Override + public void execute(Runnable task) { + task.run(); + } + }; + try { + MDC.clear(); + if (ambient) MDC.put("ambient", "saved"); + replaceAdapter((MDCAdapter) Proxy.newProxyInstance( + MDCAdapter.class.getClassLoader(), new Class[] {MDCAdapter.class}, (proxy, method, args) -> { + if (!ends.isEmpty() + && method.getName().equals(ambient ? "setContextMap" : "clear") + && injected.compareAndSet(false, true)) throw restoreFailure; + try { + return method.invoke(adapter, args); + } catch (InvocationTargetException failure) { + throw failure.getCause(); + } + })); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + if (outcome.equals("PENDING")) context.wait("pause", Duration.ofSeconds(1)); + if (outcome.equals("RETRYING")) + context.step("retry", String.class, step -> { + throw original; + }); + return "done"; + }, + DurableConfig.builder() + .withExecutorService(executor) + .withPlugins(new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + } + }) + .build()); + if (expectedFatal != null) assertSame(expectedFatal, assertThrows(Error.class, () -> runner.run("input"))); + else if (outcome.equals("RETRYING")) + assertSame( + original, + assertThrows(UnrecoverableDurableExecutionException.class, () -> runner.run("input"))); + else + assertEquals( + ExecutionStatus.valueOf(outcome), runner.run("input").getStatus()); + assertTrue(injected.get()); + assertEquals(1, ends.size()); + assertEquals(InvocationStatus.valueOf(outcome), ends.get(0).invocationStatus()); + } finally { + executor.shutdownNow(); + replaceAdapter(adapter); + if (before == null) MDC.clear(); + else MDC.setContextMap(before); + } + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void deliveryFailureRemainsRetryingUntilAReplayActuallySucceeds(boolean executionView) { + var exporter = InMemorySpanExporter.create(); + var provider = SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)); + var options = OtelPluginConfig.builder() + .contextExtractor(() -> null) + .enableMdc(false) + .build(); + DurableExecutionPlugin otel = executionView + ? new ExecutionOtelPlugin(provider, options) + : new InvocationOtelPlugin(provider, options); + var ends = new CopyOnWriteArrayList(); + var failDelivery = new AtomicBoolean(true); + var original = new IllegalStateException("cannot deliver result yet"); + var completedSideEffects = new AtomicInteger(); + var serDes = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + @Override + public String serialize(Object value) { + if ("done".equals(value) && failDelivery.compareAndSet(true, false)) throw original; + return delegate.serialize(value); + } + + @Override + public T deserialize(String value, TypeToken type) { + return delegate.deserialize(value, type); + } + }; + var runner = LocalDurableTestRunner.create( + String.class, + (input, ctx) -> { + ctx.step("saved", String.class, step -> { + completedSideEffects.incrementAndGet(); + return "checkpointed"; + }); + return "done"; + }, + DurableConfig.builder() + .withSerDes(serDes) + .withPlugins(otel, new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + } + }) + .build()); + assertSame(original, assertThrows(IllegalStateException.class, () -> runner.run("input"))); + assertEquals(InvocationStatus.RETRYING, ends.get(0).invocationStatus()); + assertSame(original, ends.get(0).executionError()); + assertEquals( + 0, + exporter.getFinishedSpanItems().stream() + .filter(s -> s.getName().equals("Workflow")) + .count()); + assertEquals(ExecutionStatus.SUCCEEDED, runner.runUntilComplete("input").getStatus()); + assertEquals( + List.of(InvocationStatus.RETRYING, InvocationStatus.SUCCEEDED), + ends.stream().map(InvocationEndInfo::invocationStatus).toList()); + assertEquals(1, completedSideEffects.get()); + assertEquals( + 1, + exporter.getFinishedSpanItems().stream() + .filter(s -> s.getName().equals("Workflow")) + .count()); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void failureSerializationRemainsRetryingUntilFailureCanBeDelivered(boolean executionView) { + var exporter = InMemorySpanExporter.create(); + var provider = SdkTracerProvider.builder().addSpanProcessor(SimpleSpanProcessor.create(exporter)); + var options = OtelPluginConfig.builder() + .contextExtractor(() -> null) + .enableMdc(false) + .build(); + DurableExecutionPlugin otel = executionView + ? new ExecutionOtelPlugin(provider, options) + : new InvocationOtelPlugin(provider, options); + var ends = new CopyOnWriteArrayList(); + var failDelivery = new AtomicBoolean(true); + var deliveryFailure = new IllegalStateException("cannot serialize failure yet"); + var bodyFailure = new IllegalArgumentException("handler failed"); + var completedSideEffects = new AtomicInteger(); + var serDes = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + @Override + public String serialize(Object value) { + if (value == bodyFailure && failDelivery.compareAndSet(true, false)) throw deliveryFailure; + return delegate.serialize(value); + } + + @Override + public T deserialize(String value, TypeToken type) { + return delegate.deserialize(value, type); + } + }; + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + context.step("saved-before-failure", String.class, step -> { + completedSideEffects.incrementAndGet(); + return "checkpointed"; + }); + throw bodyFailure; + }, + DurableConfig.builder() + .withSerDes(serDes) + .withPlugins(otel, new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + } + }) + .build()); + assertSame(deliveryFailure, assertThrows(IllegalStateException.class, () -> runner.run("input"))); + assertEquals(1, ends.size()); + assertEquals(InvocationStatus.RETRYING, ends.get(0).invocationStatus()); + assertSame(deliveryFailure, ends.get(0).executionError()); + assertEquals( + 0, + exporter.getFinishedSpanItems().stream() + .filter(span -> span.getName().equals("Workflow")) + .count()); + var firstInvocation = exporter.getFinishedSpanItems().stream() + .filter(span -> span.getName().equals("Invocation")) + .findFirst() + .orElseThrow(); + assertEquals("RETRYING", firstInvocation.getAttributes().get(DURABLE_INVOCATION_STATUS)); + assertEquals(StatusCode.UNSET, firstInvocation.getStatus().getStatusCode()); + + var result = runner.run("input"); + assertEquals(ExecutionStatus.FAILED, result.getStatus()); + assertEquals( + bodyFailure.getClass().getName(), + result.getError().orElseThrow().errorType()); + assertEquals(bodyFailure.getMessage(), result.getError().orElseThrow().errorMessage()); + assertEquals( + List.of(InvocationStatus.RETRYING, InvocationStatus.FAILED), + ends.stream().map(InvocationEndInfo::invocationStatus).toList()); + assertSame(bodyFailure, ends.get(1).executionError()); + assertEquals(1, completedSideEffects.get(), "the completed step body must not run again after delivery retry"); + assertEquals( + 1, + exporter.getFinishedSpanItems().stream() + .filter(span -> span.getName().equals("Workflow")) + .count()); + } + + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({ + "ordinary,assertion,false,false", "ordinary,assertion,true,false", + "vm,assertion,false,false", "vm,assertion,true,false", + "death,assertion,false,false", "death,assertion,true,false", + "ordinary,vm,false,false", "ordinary,vm,true,false", + "vm,death,false,false", "vm,death,true,false", + "wrapped,assertion,false,false", "wrapped,assertion,true,false", + "ordinary,assertion,false,true", "ordinary,assertion,true,true", + "vm,assertion,false,true", "vm,assertion,true,true", + "death,assertion,false,true", "death,assertion,true,true", + "ordinary,vm,false,true", "ordinary,vm,true,true", + "vm,death,false,true", "vm,death,true,true", + "wrapped,assertion,false,true", "wrapped,assertion,true,true" + }) + void preparationAndEndFailuresRetainPrimaryAndFatalIdentity( + String preparationKind, String endKind, boolean serializeFailure, boolean inline) throws Exception { + Throwable original = + switch (preparationKind) { + case "vm", "wrapped" -> new InternalError("preparation fatal"); + case "death" -> new ThreadDeath(); + default -> new IllegalStateException("preparation failure"); + }; + var delivery = preparationKind.equals("wrapped") ? new CompletionException(original) : original; + Error endFailure = + switch (endKind) { + case "vm" -> new InternalError("end fatal"); + case "death" -> new ThreadDeath(); + default -> new AssertionError("end failure"); + }; + var expected = preparationKind.equals("ordinary") && endKind.equals("vm") ? endFailure : original; + var suppressed = expected == endFailure ? original : endFailure; + var ends = new CopyOnWriteArrayList(); + var bodyFailure = new IllegalArgumentException("body failure"); + var escaped = new CountDownLatch(1); + var workerFailure = new AtomicReference(); + var callers = Executors.newSingleThreadExecutor(); + var workers = + new ThreadPoolExecutor(1, 1, 0, TimeUnit.SECONDS, new SynchronousQueue(), task -> { + var thread = new Thread(task, "combined-failure-owner"); + thread.setUncaughtExceptionHandler((owner, failure) -> { + workerFailure.set(failure); + escaped.countDown(); + }); + return thread; + }) { + @Override + public void execute(Runnable task) { + if (inline) task.run(); + else super.execute(task); + } + }; + var serDes = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + @Override + public String serialize(Object value) { + if (serializeFailure ? value == bodyFailure : "done".equals(value)) { + ExceptionHelper.sneakyThrow(delivery); + } + return delegate.serialize(value); + } + + @Override + public T deserialize(String value, TypeToken type) { + return delegate.deserialize(value, type); + } + }; + try { + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + if (serializeFailure) throw bodyFailure; + return "done"; + }, + DurableConfig.builder() + .withExecutorService(workers) + .withSerDes(serDes) + .withPlugins(new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + throw endFailure; + } + }) + .build()); + var response = callers.submit(() -> runner.run("input")); + var observed = assertThrows(ExecutionException.class, () -> response.get(3, TimeUnit.SECONDS)) + .getCause(); + assertSame(expected, observed); + assertEquals(List.of(suppressed), List.of(observed.getSuppressed())); + assertEquals(1, ends.size()); + assertEquals(InvocationStatus.RETRYING, ends.get(0).invocationStatus()); + assertSame(original, ends.get(0).executionError()); + if (!inline && (expected instanceof VirtualMachineError || expected instanceof ThreadDeath)) { + assertTrue(escaped.await(3, TimeUnit.SECONDS)); + assertSame(expected, workerFailure.get()); + } + } finally { + callers.shutdownNow(); + workers.shutdownNow(); + assertTrue(callers.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + void partialStartStillUnwindsAllConfiguredEndHooksAndRestoresBothThreads() throws Exception { + var key = ContextKey.named("partial-start-context"); + var order = new CopyOnWriteArrayList(); + var scopes = new Scope[2]; + var bodyCalls = new AtomicInteger(); + var workers = Executors.newSingleThreadExecutor(task -> new Thread( + () -> { + try (var ignored = + Context.root().with(key, "worker ambient").makeCurrent()) { + task.run(); + } + }, + "partial-start-owner")); + var first = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + order.add("start:first"); + scopes[0] = Context.current().with(key, "first").makeCurrent(); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + order.add("end:first"); + scopes[0].close(); + } + }; + var failing = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + order.add("start:failing"); + scopes[1] = Context.current().with(key, "failing").makeCurrent(); + throw new AssertionError("start failed after acquiring context"); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + order.add("end:failing"); + scopes[1].close(); + } + }; + var unstarted = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + order.add("unexpected start"); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + order.add("end:unstarted"); + } + }; + try (var ignored = Context.current().with(key, "caller ambient").makeCurrent()) { + var before = Context.current(); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + bodyCalls.incrementAndGet(); + return "done"; + }, + DurableConfig.builder() + .withExecutorService(workers) + .withPlugins(first, failing, unstarted) + .build()); + assertEquals(ExecutionStatus.FAILED, runner.run("input").getStatus()); + assertEquals(List.of("start:first", "start:failing", "end:unstarted", "end:failing", "end:first"), order); + assertEquals(0, bodyCalls.get()); + assertSame(before, Context.current()); + assertEquals( + "worker ambient", + workers.submit(() -> Context.current().get(key)).get(5, TimeUnit.SECONDS)); + } finally { + workers.shutdownNow(); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + private static UnrecoverableDurableExecutionException retryError() { + return new UnrecoverableDurableExecutionException( + ErrorObject.builder().errorMessage("retry").build(), true); + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(5, TimeUnit.SECONDS)); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + + private static void replaceAdapter(MDCAdapter adapter) throws Exception { + var setter = MDC.class.getDeclaredMethod("setMDCAdapter", MDCAdapter.class); + setter.setAccessible(true); + setter.invoke(null, adapter); + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/LegacySubclassScopeCompatibilityTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/LegacySubclassScopeCompatibilityTest.java new file mode 100644 index 000000000..2926a3506 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/LegacySubclassScopeCompatibilityTest.java @@ -0,0 +1,91 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import java.net.URL; +import java.net.URLClassLoader; +import java.nio.file.Files; +import java.nio.file.Path; +import java.time.Instant; +import javax.tools.ToolProvider; +import org.junit.jupiter.api.io.TempDir; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; + +class LegacySubclassScopeCompatibilityTest { + @TempDir + Path directory; + + @ParameterizedTest + @CsvSource({"Invocation,false", "Invocation,true", "Execution,false", "Execution,true"}) + void preservesOldSubclassMethodsAndUnrelatedDefaults(String view, boolean ownMethod) throws Exception { + var baseline = Files.createDirectories(directory.resolve("baseline")); + var classes = Files.createDirectories(directory.resolve("classes")); + var baseName = view + "OtelPlugin"; + var oldBase = directory.resolve(baseName + ".java"); + Files.writeString(oldBase, "package software.amazon.lambda.durable.otel; public class " + baseName + " {}"); + var compiler = ToolProvider.getSystemJavaCompiler(); + assertEquals( + 0, compiler.run(null, null, null, "--release", "17", "-d", baseline.toString(), oldBase.toString())); + var source = directory.resolve("LegacySubclass.java"); + var method = "public AutoCloseable openHandlerScope() { opened++; return () -> closed++; }"; + Files.writeString(source, """ + import software.amazon.lambda.durable.otel.%s; + interface ApplicationScope { + default AutoCloseable openHandlerScope() { + LegacySubclass.opened++; + return () -> LegacySubclass.closed++; + } + } + public class LegacySubclass extends %s implements ApplicationScope { + public static int opened, closed; + %s + public void originalCall() throws Exception { try (var scope = openHandlerScope()) {} } + } + """.formatted(baseName, baseName, ownMethod ? method : "")); + assertEquals( + 0, + compiler.run( + null, + null, + null, + "--release", + "17", + "-cp", + baseline.toString(), + "-d", + classes.toString(), + source.toString())); + try (var loader = new URLClassLoader( + new URL[] {classes.toUri().toURL()}, getClass().getClassLoader())) { + var type = loader.loadClass("LegacySubclass"); + var plugin = (DurableExecutionPlugin) type.getConstructor().newInstance(); + plugin.onInvocationStart(new InvocationInfo("req", "arn", true, Instant.now())); + plugin.onInvocationEnd(new InvocationEndInfo("req", "arn", true, InvocationStatus.SUCCEEDED, null)); + assertEquals(0, type.getField("opened").get(null), "SDK must not invoke an unrelated subclass resource"); + type.getMethod("originalCall").invoke(plugin); + assertEquals( + 1, type.getField("opened").get(null), "new SDK members must not shadow an application default"); + assertEquals(1, type.getField("closed").get(null)); + } + assertEquals( + 0, + compiler.run( + null, + null, + null, + "--release", + "17", + "-cp", + System.getProperty("java.class.path"), + "-d", + classes.toString(), + source.toString())); + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcFailureBoundaryTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcFailureBoundaryTest.java new file mode 100644 index 000000000..93cce8a8a --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcFailureBoundaryTest.java @@ -0,0 +1,446 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Proxy; +import java.lang.reflect.UndeclaredThrowableException; +import java.time.Duration; +import java.util.List; +import java.util.concurrent.AbstractExecutorService; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.stream.Stream; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.MethodSource; +import org.slf4j.MDC; +import org.slf4j.spi.MDCAdapter; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; + +class MdcFailureBoundaryTest { + @Test + void workerCaptureFailureDoesNotEndAnInvocationThatNeverStarted() throws Exception { + var original = MDC.getMDCAdapter(); + var injected = new AtomicBoolean(); + var failure = new IllegalStateException("worker MDC capture failed"); + var plugin = new RecordingPlugin(); + var bodyCalls = new AtomicInteger(); + var executor = new CompletingWorker(); + try { + replaceAdapter(proxy(original, (method) -> { + if (method.equals("getCopyOfContextMap") + && Thread.currentThread().getName().equals("mdc-initialization-worker") + && injected.compareAndSet(false, true)) throw failure; + return null; + })); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + bodyCalls.incrementAndGet(); + return "done"; + }, + DurableConfig.builder() + .withExecutorService(executor) + .withPlugins(plugin) + .build()); + var result = runner.run("input"); + assertTrue(injected.get()); + assertEquals(ExecutionStatus.FAILED, result.getStatus()); + assertEquals(failure.getMessage(), result.getError().orElseThrow().errorMessage()); + assertEquals(0, bodyCalls.get()); + assertEquals(0, plugin.starts.get()); + assertEquals(0, plugin.ends.get(), "no end hook may consume stale state before start dispatch"); + } finally { + try { + executor.shutdownNow(); + assertTrue(executor.awaitTermination(3, TimeUnit.SECONDS)); + } finally { + replaceAdapter(original); + } + } + } + + @ParameterizedTest + @CsvSource({ + "false,false,exception", "true,false,exception", "false,true,exception", "true,true,exception", + "false,false,fatal", "true,false,fatal", "false,true,fatal", "true,true,fatal", + "false,false,assertion", "true,false,assertion", "false,true,assertion", "true,true,assertion", + "false,false,linkage", "true,false,linkage", "false,true,linkage", "true,true,linkage" + }) + void workerMdcRestorationFailureCannotStrandTheInvocation(boolean ambientMdc, boolean suspend, String failureKind) + throws Exception { + var original = MDC.getMDCAdapter(); + var injected = new AtomicBoolean(); + var plugin = new RecordingPlugin(); + Throwable failure = + switch (failureKind) { + case "fatal" -> new InternalError("worker MDC restoration failed"); + case "assertion" -> new AssertionError("worker MDC restoration failed"); + case "linkage" -> new NoClassDefFoundError("worker MDC restoration failed"); + default -> new IllegalStateException("worker MDC restoration failed"); + }; + var ownerFailure = new AtomicReference(); + var observed = new CountDownLatch(1); + var workers = Executors.newSingleThreadExecutor(task -> { + var thread = new Thread( + () -> { + if (ambientMdc) MDC.put("ambient", "saved"); + else MDC.clear(); + task.run(); + }, + "mdc-restoration-worker"); + thread.setUncaughtExceptionHandler((owner, error) -> { + ownerFailure.set(error); + observed.countDown(); + }); + return thread; + }); + try { + replaceAdapter(proxy(original, method -> { + if (Thread.currentThread().getName().equals("mdc-restoration-worker") + && plugin.ends.get() == 1 + && method.equals(ambientMdc ? "setContextMap" : "clear") + && injected.compareAndSet(false, true)) throw failure; + return null; + })); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + if (suspend) context.wait("resume", Duration.ofSeconds(1)); + return "done"; + }, + DurableConfig.builder() + .withExecutorService(workers) + .withPlugins(plugin) + .build()); + if (failureKind.equals("fatal")) { + assertSame(failure, assertThrows(Error.class, () -> runner.run("input"))); + } else { + var result = runner.run("input"); + assertEquals(suspend ? ExecutionStatus.PENDING : ExecutionStatus.SUCCEEDED, result.getStatus()); + } + assertEquals(1, plugin.starts.get()); + assertEquals(1, plugin.ends.get()); + assertTrue(injected.get()); + assertTrue(observed.await(3, TimeUnit.SECONDS)); + assertSame(failure, ownerFailure.get(), "restoration failure still escapes its worker after End"); + } finally { + workers.shutdownNow(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + replaceAdapter(original); + } + } + + static Stream postEndFatalCases() { + return Stream.of(false, true) + .flatMap(inline -> Stream.of(false, true) + .flatMap(suspend -> Stream.of(false, true) + .flatMap(death -> Stream.of("direct", "completion", "execution", "reflection", "proxy") + .map(wrapper -> Arguments.of(inline, suspend, death, wrapper))))); + } + + @SuppressWarnings("removal") + @ParameterizedTest + @MethodSource("postEndFatalCases") + void postEndFatalReachesCallerAndWorkerWithoutChangingEnd( + boolean inline, boolean suspend, boolean death, String wrapper) throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("post-End restore fatal"); + Throwable failure = + switch (wrapper) { + case "completion" -> new CompletionException(fatal); + case "execution" -> new ExecutionException(fatal); + case "reflection" -> new InvocationTargetException(fatal); + case "proxy" -> new UndeclaredThrowableException(fatal); + default -> fatal; + }; + exercisePostEndFailure(inline, suspend, failure, fatal, null); + } + + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({ + "false,false,false", + "false,false,true", + "false,true,false", + "false,true,true", + "true,false,false", + "true,false,true", + "true,true,false", + "true,true,true" + }) + void fatalDiagnosticRetainsIdentityAndIsReadOnce(boolean inline, boolean suspend, boolean death) throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("fatal diagnostic"); + var reads = new AtomicInteger(); + var wrapper = new CompletionException("diagnostic", null) { + @Override + public synchronized Throwable getCause() { + reads.incrementAndGet(); + throw fatal; + } + }; + exercisePostEndFailure(inline, suspend, wrapper, fatal, null); + assertEquals(1, reads.get()); + } + + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void earlierEndFatalRemainsPrimaryAndRetainsRestorationDiagnostic(boolean inline, boolean death) throws Exception { + Error primary = death ? new ThreadDeath() : new InternalError("End fatal"); + var restoration = new InternalError("later restoration fatal"); + exercisePostEndFailure(inline, false, restoration, primary, primary); + assertArrayEquals(new Throwable[] {restoration}, primary.getSuppressed()); + } + + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({ + "false,false,false", + "false,false,true", + "false,true,false", + "false,true,true", + "true,false,false", + "true,false,true", + "true,true,false", + "true,true,true" + }) + void sameFatalAtEndAndRestorationRetainsOwnerIdentity(boolean inline, boolean suspend, boolean death) + throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("same End/restore fatal"); + exercisePostEndFailure(inline, suspend, fatal, fatal, fatal); + assertEquals(0, fatal.getSuppressed().length, "Never suppress the fatal onto itself"); + } + + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void restorationFatalRetainsEarlierNonfatalEndFailure(boolean inline, boolean wrapped) throws Exception { + var earlier = new AssertionError("ordinary End failure"); + var fatal = new InternalError("later restoration fatal"); + exercisePostEndFailure(inline, false, wrapped ? new CompletionException(fatal) : fatal, fatal, earlier); + assertArrayEquals(new Throwable[] {earlier}, fatal.getSuppressed()); + } + + private static void exercisePostEndFailure( + boolean inline, boolean suspend, Throwable failure, Error expectedFatal, Error endFailure) + throws Exception { + var original = MDC.getMDCAdapter(); + var injected = new AtomicBoolean(); + var plugin = new RecordingPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + super.onInvocationEnd(info); + if (endFailure != null) throw endFailure; + } + }; + var ownerFailure = new AtomicReference(); + var owner = new AtomicReference(); + var escaped = new CountDownLatch(1); + var caller = Executors.newSingleThreadExecutor(task -> new Thread(task, "post-end-caller")); + var delegate = Executors.newSingleThreadExecutor(task -> { + var thread = new Thread(task, "post-end-worker"); + thread.setUncaughtExceptionHandler((worker, error) -> { + ownerFailure.set(error); + escaped.countDown(); + }); + return thread; + }); + var executor = new AbstractExecutorService() { + @Override + public void execute(Runnable task) { + Runnable owned = () -> { + owner.set(Thread.currentThread()); + if (suspend) MDC.put("ambient", "saved"); + else MDC.clear(); + task.run(); + }; + if (inline) { + try { + owned.run(); + } catch (Throwable error) { + ownerFailure.set(error); + escaped.countDown(); + throw error; + } + } else delegate.execute(owned); + } + + @Override + public void shutdown() { + delegate.shutdown(); + } + + @Override + public List shutdownNow() { + return delegate.shutdownNow(); + } + + @Override + public boolean isShutdown() { + return delegate.isShutdown(); + } + + @Override + public boolean isTerminated() { + return delegate.isTerminated(); + } + + @Override + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(timeout, unit); + } + }; + try { + replaceAdapter(proxy(original, method -> { + if (Thread.currentThread() == owner.get() + && plugin.ends.get() == 1 + && method.equals(suspend ? "setContextMap" : "clear") + && injected.compareAndSet(false, true)) throw failure; + return null; + })); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + if (suspend) context.wait("resume", Duration.ofSeconds(1)); + return "done"; + }, + DurableConfig.builder() + .withExecutorService(executor) + .withPlugins(plugin) + .build()); + var observation = caller.submit(() -> { + try { + return (Object) runner.run("input"); + } catch (Throwable error) { + return error; + } + }); + assertSame( + expectedFatal, + observation.get(3, TimeUnit.SECONDS), + "caller receives original fatal without hanging"); + assertTrue(escaped.await(3, TimeUnit.SECONDS), "fatal escapes actual worker after observation settlement"); + assertSame(expectedFatal, ownerFailure.get()); + assertTrue(injected.get()); + assertEquals(1, plugin.starts.get()); + assertEquals(1, plugin.ends.get()); + assertEquals(suspend ? InvocationStatus.PENDING : InvocationStatus.SUCCEEDED, plugin.status.get()); + } finally { + caller.shutdownNow(); + executor.shutdownNow(); + try { + assertTrue(caller.awaitTermination(3, TimeUnit.SECONDS)); + assertTrue(executor.awaitTermination(3, TimeUnit.SECONDS)); + } finally { + replaceAdapter(original); + } + } + } + + @FunctionalInterface + private interface Fault { + Object apply(String method) throws Throwable; + } + + private static MDCAdapter proxy(MDCAdapter delegate, Fault fault) { + return (MDCAdapter) Proxy.newProxyInstance( + MDCAdapter.class.getClassLoader(), new Class[] {MDCAdapter.class}, (proxy, method, args) -> { + var value = fault.apply(method.getName()); + if (value != null) return value; + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException ex) { + throw ex.getCause(); + } + }); + } + + private static void replaceAdapter(MDCAdapter adapter) throws Exception { + var setter = MDC.class.getDeclaredMethod("setMDCAdapter", MDCAdapter.class); + setter.setAccessible(true); + setter.invoke(null, adapter); + } + + private static class RecordingPlugin implements DurableExecutionPlugin { + final AtomicInteger starts = new AtomicInteger(); + final AtomicInteger ends = new AtomicInteger(); + final AtomicReference status = new AtomicReference<>(); + + @Override + public void onInvocationStart(InvocationInfo info) { + starts.incrementAndGet(); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + status.set(info.invocationStatus()); + ends.incrementAndGet(); + } + } + /** Executes on a real worker, completing its first task before returning to the invocation caller. */ + private static final class CompletingWorker extends AbstractExecutorService { + private final ExecutorService delegate = + Executors.newSingleThreadExecutor(task -> new Thread(task, "mdc-initialization-worker")); + private final AtomicBoolean first = new AtomicBoolean(true); + + @Override + public void execute(Runnable task) { + if (!first.compareAndSet(true, false)) { + delegate.execute(task); + return; + } + var done = new CompletableFuture(); + delegate.execute(() -> { + try { + task.run(); + } finally { + done.complete(null); + } + }); + done.orTimeout(3, TimeUnit.SECONDS).join(); + } + + @Override + public void shutdown() { + delegate.shutdown(); + } + + @Override + public List shutdownNow() { + return delegate.shutdownNow(); + } + + @Override + public boolean isShutdown() { + return delegate.isShutdown(); + } + + @Override + public boolean isTerminated() { + return delegate.isTerminated(); + } + + @Override + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(timeout, unit); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcFatalInitializationTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcFatalInitializationTest.java new file mode 100644 index 000000000..615a4be86 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcFatalInitializationTest.java @@ -0,0 +1,383 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Proxy; +import java.lang.reflect.UndeclaredThrowableException; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.AbstractExecutorService; +import java.util.concurrent.CompletionException; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Supplier; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import org.slf4j.MDC; +import org.slf4j.spi.MDCAdapter; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; + +class MdcFatalInitializationTest { + @ParameterizedTest + @MethodSource("fatalCases") + @SuppressWarnings("removal") + void fatalCaptureSettlesCallerThenEscapesOwnerBeforeAnyHook(String mode, boolean death, String wrapper) + throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("fatal capture"); + Throwable failure = + switch (wrapper) { + case "completion" -> new CompletionException(fatal); + case "nested" -> new CompletionException(new ExecutionException(fatal)); + case "reflection" -> new InvocationTargetException(fatal); + case "proxy" -> new UndeclaredThrowableException(fatal); + default -> fatal; + }; + runCapture(mode, failure, fatal); + } + + @ParameterizedTest + @MethodSource("ordinaryCases") + void nonfatalInitializationKeepsItsDurableFailurePolicy(String mode, String kind) throws Exception { + Throwable failure = + switch (kind) { + case "assertion" -> new AssertionError("ordinary capture"); + case "application-cause" -> + new IllegalStateException("ordinary capture", new InternalError("diagnostic")); + default -> new IllegalStateException("ordinary capture"); + }; + runCapture(mode, failure, null); + } + + @ParameterizedTest + @MethodSource("diagnosticCases") + @SuppressWarnings("removal") + void captureDiagnosticsAreBoundedAndPreserveTheSelectedFailure(String mode, String kind) throws Exception { + var abort = new AtomicBoolean(); + var reads = new ArrayList(); + var leaf = new IllegalStateException("ordinary capture"); + Error fatal = kind.equals("getter-death") ? new ThreadDeath() : new InternalError("fatal capture diagnostic"); + Throwable failure; + Throwable expected; + if (kind.equals("self")) { + var link = new AtomicReference(); + failure = diagnostic(link::get, abort, reads); + link.set(failure); + expected = failure; + } else if (kind.equals("pair")) { + var link = new AtomicReference(); + var inner = diagnostic(link::get, abort, reads); + failure = diagnostic(() -> inner, abort, reads); + link.set(failure); + expected = failure; + } else if (kind.equals("null")) { + failure = diagnostic(() -> null, abort, reads); + expected = failure; + } else if (kind.equals("unreadable")) { + failure = diagnostic( + () -> { + throw new IllegalStateException("unreadable cause"); + }, + abort, + reads); + expected = failure; + } else if (kind.equals("nested")) { + var inner = diagnostic(() -> leaf, abort, reads); + failure = diagnostic(() -> inner, abort, reads); + expected = leaf; + } else if (kind.equals("non-completion")) { + var wrapper = new UndeclaredThrowableException(leaf, "ordinary capture"); + failure = diagnostic(() -> wrapper, abort, reads); + expected = wrapper; // Ordinary initialization only unwraps a leading CompletionException chain. + } else if (kind.equals("changing")) { + var calls = new AtomicInteger(); + failure = diagnostic( + () -> { + if (calls.incrementAndGet() == 1) return leaf; + throw fatal; + }, + abort, + reads); + expected = leaf; + } else { + failure = diagnostic( + () -> { + throw fatal; + }, + abort, + reads); + expected = failure; + } + var serialized = new AtomicReference(); + var safeSerDes = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + @Override + public String serialize(Object value) { + if (value instanceof Throwable error) { + serialized.set(error); + // Observe identity without making a customer serializer traverse the malformed cause graph. + return "\"capture failure\""; + } + return delegate.serialize(value); + } + + @Override + public T deserialize(String value, TypeToken type) { + return delegate.deserialize(value, type); + } + }; + var getterFatal = kind.startsWith("getter-"); + runCapture(mode, failure, getterFatal ? fatal : null, expected, safeSerDes, () -> abort.set(true)); + if (getterFatal) assertNull(serialized.get()); + else assertSame(expected, serialized.get(), "initialization must preserve the selected failure identity"); + reads.forEach(count -> assertEquals(1, count.get(), "read each diagnostic cause once")); + } + + private static CompletionException diagnostic( + Supplier cause, AtomicBoolean abort, List reads) { + var count = new AtomicInteger(); + reads.add(count); + return new CompletionException("ordinary capture", null) { + @Override + public synchronized Throwable getCause() { + if (abort.get()) return null; // Bound negative-control cleanup even when old code loops forever. + count.incrementAndGet(); + return cause.get(); + } + }; + } + + private static Stream diagnosticCases() { + return Stream.of("direct", "async") + .flatMap(mode -> Stream.of( + "self", + "pair", + "null", + "unreadable", + "nested", + "non-completion", + "changing", + "getter-vm", + "getter-death") + .map(kind -> Arguments.of(mode, kind))); + } + + private static Stream fatalCases() { + return Stream.of("direct", "async", "precompleted") + .flatMap(mode -> Stream.of(false, true) + .flatMap(death -> Stream.of("direct", "completion", "nested", "reflection", "proxy") + .map(wrapper -> Arguments.of(mode, death, wrapper)))); + } + + private static Stream ordinaryCases() { + return Stream.of("direct", "async", "precompleted") + .flatMap(mode -> + Stream.of("exception", "assertion", "application-cause").map(kind -> Arguments.of(mode, kind))); + } + + private static void runCapture(String mode, Throwable failure, Error expectedFatal) throws Exception { + runCapture(mode, failure, expectedFatal, failure, null, () -> {}); + } + + private static void runCapture( + String mode, + Throwable failure, + Error expectedFatal, + Throwable expectedOrdinary, + SerDes overrideSerDes, + Runnable release) + throws Exception { + var originalAdapter = MDC.getMDCAdapter(); + var injected = new AtomicBoolean(); + var starts = new AtomicInteger(); + var bodies = new AtomicInteger(); + var ends = new AtomicInteger(); + var callerThread = new AtomicReference(); + var workers = new ObservedExecutor(mode); + var caller = Executors.newSingleThreadExecutor(task -> daemon(task, "capture-caller")); + try { + replaceAdapter((MDCAdapter) Proxy.newProxyInstance( + MDCAdapter.class.getClassLoader(), new Class[] {MDCAdapter.class}, (proxy, method, args) -> { + var capture = method.getName().equals("getCopyOfContextMap") + && Arrays.stream(Thread.currentThread().getStackTrace()) + .anyMatch(frame -> frame.getClassName().endsWith(".DurableExecutor") + && frame.getMethodName().equals("restoreMdcOnClose")); + if (capture && injected.compareAndSet(false, true)) throw failure; + try { + return method.invoke(originalAdapter, args); + } catch (InvocationTargetException invocation) { + throw invocation.getCause(); + } + })); + var plugin = new DurableExecutionPlugin() { + @Override + public void onInvocationStart(InvocationInfo info) { + starts.incrementAndGet(); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.incrementAndGet(); + } + }; + var configuration = + DurableConfig.builder().withExecutorService(workers).withPlugins(plugin); + if (overrideSerDes != null) configuration.withSerDes(overrideSerDes); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + bodies.incrementAndGet(); + return "unreachable"; + }, + configuration.build()); + var response = caller.submit(() -> { + callerThread.set(Thread.currentThread()); + return runner.run("input"); + }); + if (expectedFatal != null) { + var thrown = assertThrows( + ExecutionException.class, + () -> response.get(3, TimeUnit.SECONDS), + "the observation future must settle rather than strand the invocation"); + assertSame(expectedFatal, thrown.getCause()); + assertTrue(workers.escaped.await(3, TimeUnit.SECONDS)); + assertSame(expectedFatal, workers.escape.get()); + if (mode.equals("direct")) assertSame(callerThread.get(), workers.owner.get()); + else { + assertNotSame(callerThread.get(), workers.owner.get()); + assertTrue(workers.uncaught.await(3, TimeUnit.SECONDS)); + assertSame(expectedFatal, workers.uncaughtFailure.get()); + } + } else { + var result = response.get(3, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.FAILED, result.getStatus()); + assertEquals( + expectedOrdinary.getClass().getName(), + result.getError().orElseThrow().errorType()); + assertEquals("ordinary capture", result.getError().orElseThrow().errorMessage()); + assertNull(workers.escape.get(), "ordinary initialization policy must not be broadened"); + } + assertTrue(injected.get()); + assertEquals(0, starts.get()); + assertEquals(0, bodies.get()); + assertEquals(0, ends.get()); + } finally { + release.run(); + workers.shutdownNow(); + caller.shutdownNow(); + try { + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + assertTrue(caller.awaitTermination(3, TimeUnit.SECONDS)); + } finally { + replaceAdapter(originalAdapter); + } + } + } + + private static Thread daemon(Runnable task, String name) { + var thread = new Thread(task, name); + thread.setDaemon(true); + return thread; + } + + private static void replaceAdapter(MDCAdapter adapter) throws Exception { + var setter = MDC.class.getDeclaredMethod("setMDCAdapter", MDCAdapter.class); + setter.setAccessible(true); + setter.invoke(null, adapter); + } + + private static final class ObservedExecutor extends AbstractExecutorService { + private final String mode; + private final AtomicReference escape = new AtomicReference<>(); + private final AtomicReference owner = new AtomicReference<>(); + private final CountDownLatch escaped = new CountDownLatch(1); + private final AtomicReference uncaughtFailure = new AtomicReference<>(); + private final CountDownLatch uncaught = new CountDownLatch(1); + private final ExecutorService delegate = Executors.newSingleThreadExecutor(task -> { + var thread = daemon(task, "capture-worker"); + thread.setUncaughtExceptionHandler((worker, failure) -> { + uncaughtFailure.set(failure); + uncaught.countDown(); + }); + return thread; + }); + + private ObservedExecutor(String mode) { + this.mode = mode; + } + + @Override + public void execute(Runnable task) { + var completed = new CountDownLatch(1); + Runnable observed = () -> { + owner.set(Thread.currentThread()); + try { + task.run(); + } catch (Error failure) { + escape.set(failure); + escaped.countDown(); + throw failure; + } finally { + completed.countDown(); + } + }; + if (mode.equals("direct")) observed.run(); + else { + delegate.execute(observed); + if (mode.equals("precompleted")) { + try { + assertTrue(completed.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + } + } + + @Override + public void shutdown() { + delegate.shutdown(); + } + + @Override + public List shutdownNow() { + return delegate.shutdownNow(); + } + + @Override + public boolean isShutdown() { + return delegate.isShutdown(); + } + + @Override + public boolean isTerminated() { + return delegate.isTerminated(); + } + + @Override + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(timeout, unit); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcWorkerReuseBoundaryTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcWorkerReuseBoundaryTest.java new file mode 100644 index 000000000..60d7224c6 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/MdcWorkerReuseBoundaryTest.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.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Proxy; +import java.time.Duration; +import java.util.List; +import java.util.concurrent.*; +import java.util.concurrent.atomic.*; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import org.junit.jupiter.params.provider.ValueSource; +import org.slf4j.MDC; +import org.slf4j.helpers.BasicMDCAdapter; +import org.slf4j.spi.MDCAdapter; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.*; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; + +class MdcWorkerReuseBoundaryTest { + @ParameterizedTest + @ValueSource(strings = {"start", "end"}) + void replacementWorkerDoesNotInheritInvocationMdcWhenClearSucceeds(String markerPhase) throws Exception { + exercise(markerPhase, false, "ok"); + } + + static Stream fallbackCases() { + return Stream.of(false, true) + .flatMap(suspend -> Stream.of( + "ordinary", "assertion", "same", "vm", "death", "wrapped-vm", "wrapped-death") + .map(kind -> Arguments.of(suspend, kind))); + } + + @ParameterizedTest + @MethodSource("fallbackCases") + void failedFallbackRetainsFailurePolicyAndFatalIdentity(boolean suspend, String kind) throws Exception { + exercise("end", suspend, kind); + } + + static Stream fatalEndCases() { + return Stream.of(false, true) + .flatMap(suspend -> Stream.of(false, true) + .flatMap(death -> Stream.of( + "ok", + "ordinary", + "assertion", + "same", + "vm", + "death", + "wrapped-vm", + "wrapped-death") + .map(kind -> Arguments.of(suspend, death, kind)))); + } + + @ParameterizedTest + @MethodSource("fatalEndCases") + @SuppressWarnings("removal") + void fatalEndStillClearsOrdinaryRestoreFailureBeforeWorkerReplacement(boolean suspend, boolean death, String kind) + throws Exception { + exercise("end", suspend, kind, death ? new ThreadDeath() : new InternalError("original End fatal")); + } + + private static void exercise(String markerPhase, boolean suspend, String kind) throws Exception { + exercise(markerPhase, suspend, kind, null); + } + + @SuppressWarnings("removal") + private static void exercise(String markerPhase, boolean suspend, String kind, Error selectedFatal) + throws Exception { + var original = MDC.getMDCAdapter(); + var basic = new BasicMDCAdapter(); + var ends = new AtomicInteger(); + var endStatus = new AtomicReference(); + var failed = new AtomicBoolean(); + var clearTried = new AtomicBoolean(); + var restoreFailure = new IllegalStateException("restore failed"); + Error clearFatal = + switch (kind) { + case "vm", "wrapped-vm" -> new InternalError("fallback fatal"); + case "death", "wrapped-death" -> new ThreadDeath(); + default -> null; + }; + Throwable clearFailure = + switch (kind) { + case "ordinary" -> new IllegalArgumentException("clear failed"); + case "assertion" -> new AssertionError("clear failed"); + case "same" -> restoreFailure; + case "wrapped-vm", "wrapped-death" -> new CompletionException(new ExecutionException(clearFatal)); + default -> clearFatal; + }; + Error expectedFatal = selectedFatal != null ? selectedFatal : clearFatal; + var escaped = new CountDownLatch(1); + var failedWorker = new AtomicReference(); + var ownerFailure = new AtomicReference(); + var dirtyAtFailure = new AtomicReference(); + var count = new AtomicInteger(); + var workers = Executors.newSingleThreadExecutor(task -> { + var thread = new Thread(task, "mdc-inheritable-worker-" + count.incrementAndGet()); + thread.setUncaughtExceptionHandler((owner, failure) -> { + failedWorker.set(owner); + ownerFailure.set(failure); + dirtyAtFailure.set(MDC.get("invocation-marker")); + escaped.countDown(); + }); + return thread; + }); + try { + var adapter = (MDCAdapter) Proxy.newProxyInstance( + MDCAdapter.class.getClassLoader(), new Class[] {MDCAdapter.class}, (proxy, method, args) -> { + if (ends.get() == 1 + && method.getName().equals("setContextMap") + && Thread.currentThread().getName().startsWith("mdc-inheritable-worker-") + && failed.compareAndSet(false, true)) throw restoreFailure; + if (failed.get() + && method.getName().equals("clear") + && Thread.currentThread().getName().startsWith("mdc-inheritable-worker-") + && clearTried.compareAndSet(false, true) + && clearFailure != null) throw clearFailure; + try { + return method.invoke(basic, args); + } catch (InvocationTargetException failure) { + throw failure.getCause(); + } + }); + replace(adapter); + MDC.put("ambient", "saved"); + var plugin = new DurableExecutionPlugin() { + public void onInvocationStart(InvocationInfo info) { + if (markerPhase.equals("start")) MDC.put("invocation-marker", "previous-invocation"); + } + + public void onInvocationEnd(InvocationEndInfo info) { + if (markerPhase.equals("end")) MDC.put("invocation-marker", "previous-invocation"); + endStatus.set(info.invocationStatus()); + ends.incrementAndGet(); + if (selectedFatal != null) throw selectedFatal; + } + }; + var runner = LocalDurableTestRunner.create( + String.class, + (value, context) -> { + if (suspend) context.wait("resume", Duration.ofSeconds(1)); + return "done"; + }, + DurableConfig.builder() + .withExecutorService(workers) + .withPlugins(plugin) + .build()); + if (expectedFatal != null) assertSame(expectedFatal, assertThrows(Error.class, () -> runner.run("input"))); + else + assertEquals( + suspend ? ExecutionStatus.PENDING : ExecutionStatus.SUCCEEDED, + runner.run("input").getStatus()); + assertTrue(escaped.await(3, TimeUnit.SECONDS)); + assertSame(expectedFatal != null ? expectedFatal : restoreFailure, ownerFailure.get()); + + assertEquals(1, ends.get()); + assertEquals(suspend ? InvocationStatus.PENDING : InvocationStatus.SUCCEEDED, endStatus.get()); + if (selectedFatal != null) { + assertEquals( + clearFailure != null && clearFailure != restoreFailure + ? List.of(restoreFailure, clearFailure) + : List.of(restoreFailure), + List.of(selectedFatal.getSuppressed())); + } else if (expectedFatal != null) + assertEquals(List.of(restoreFailure), List.of(expectedFatal.getSuppressed())); + else + assertEquals( + clearFailure != null && clearFailure != restoreFailure ? List.of(clearFailure) : List.of(), + List.of(restoreFailure.getSuppressed())); + var nextThread = workers.submit(Thread::currentThread).get(3, TimeUnit.SECONDS); + var inherited = workers.submit(() -> MDC.get("invocation-marker")).get(3, TimeUnit.SECONDS); + assertNotSame(failedWorker.get(), nextThread); + System.out.println("MDC_REUSE_PUBLIC phase=" + markerPhase + " kind=" + kind + " replaced=true oldMarker=" + + dirtyAtFailure.get() + " inheritedMarker=" + inherited); + if (clearFailure == null) assertNull(inherited, "Successful fallback clears inheritable invocation state"); + assertTrue(clearTried.get()); + // An adapter whose clear also fails cannot promise clean replacement state; its failure is exposed. + } finally { + workers.shutdownNow(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + basic.clear(); + replace(original); + } + } + + private static void replace(MDCAdapter adapter) throws Exception { + var setter = MDC.class.getDeclaredMethod("setMDCAdapter", MDCAdapter.class); + setter.setAccessible(true); + setter.invoke(null, adapter); + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/ReplacementSamplerPolicyTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/ReplacementSamplerPolicyTest.java new file mode 100644 index 000000000..2d3fbe8a5 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/ReplacementSamplerPolicyTest.java @@ -0,0 +1,127 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.*; + +import io.opentelemetry.api.GlobalOpenTelemetry; +import io.opentelemetry.api.OpenTelemetry; +import io.opentelemetry.api.common.Attributes; +import io.opentelemetry.api.trace.Span; +import io.opentelemetry.api.trace.SpanKind; +import io.opentelemetry.api.trace.Tracer; +import io.opentelemetry.api.trace.TracerProvider; +import io.opentelemetry.context.Context; +import io.opentelemetry.context.propagation.ContextPropagators; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.data.LinkData; +import io.opentelemetry.sdk.trace.samplers.Sampler; +import io.opentelemetry.sdk.trace.samplers.SamplingResult; +import java.time.Duration; +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; + +class ReplacementSamplerPolicyTest { + @ParameterizedTest + @CsvSource({ + "invocation,false,root,false,plain", "execution,false,root,false,plain", + "invocation,true,root,true,plain", "execution,true,root,true,plain", + "invocation,false,sampled,true,plain", "execution,false,sampled,true,plain", + "invocation,true,unsampled,false,plain", "execution,true,unsampled,false,plain", + "invocation,false,root,false,local", "execution,false,root,false,local", + "invocation,false,root,false,opaque", "execution,false,root,false,opaque" + }) + void visibleReplacementParentBasedSamplerKeepsItsRootPolicy( + String view, boolean rootSamples, String upstream, boolean expectedSampled, String topology) { + var oldHeader = System.getProperty("com.amazonaws.xray.traceHeader"); + var evaluations = new AtomicInteger(); + var sideEffects = new AtomicInteger(); + var decisions = new CopyOnWriteArrayList(); + var header = "Root=1-6955b900-123456789012345678901234"; + if (upstream.equals("sampled")) header += ";Sampled=1"; + else if (upstream.equals("unsampled")) header += ";Sampled=0"; + var root = new Sampler() { + public SamplingResult shouldSample( + Context parent, + String traceId, + String name, + SpanKind kind, + Attributes attributes, + List links) { + evaluations.incrementAndGet(); + assertFalse(Span.fromContext(parent).getSpanContext().isValid()); + return rootSamples ? SamplingResult.recordAndSample() : SamplingResult.drop(); + } + + public String getDescription() { + return "counted-root-policy"; + } + }; + GlobalOpenTelemetry.resetForTest(); + OtelPluginAutoConfigurationState.markInstalled(); + System.setProperty("com.amazonaws.xray.traceHeader", header); + var configured = Sampler.parentBased(root); + try (var provider = SdkTracerProvider.builder() + .setSampler(topology.equals("plain") ? configured : DurableSampler.wrap(configured)) + .setIdGenerator(new DeterministicIdGenerator()) + .build()) { + var opaque = new TracerProvider() { + public Tracer get(String name) { + return provider.get(name); + } + + public Tracer get(String name, String version) { + return provider.get(name, version); + } + }; + GlobalOpenTelemetry.set(new OpenTelemetry() { + public TracerProvider getTracerProvider() { + return topology.equals("opaque") ? opaque : provider; + } + + public ContextPropagators getPropagators() { + return ContextPropagators.noop(); + } + }); + DurableExecutionPlugin plugin = + view.equals("invocation") ? new InvocationOtelPlugin() : new ExecutionOtelPlugin(); + var tracer = provider.get("application"); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + var span = tracer.spanBuilder("application-work").startSpan(); + decisions.add(span.isRecording()); + span.end(); + context.step("saved", String.class, step -> { + sideEffects.incrementAndGet(); + return "saved"; + }); + context.wait("pause", Duration.ofSeconds(1)); + return "done"; + }, + DurableConfig.builder().withPlugins(plugin).build()); + assertEquals(ExecutionStatus.PENDING, runner.run("input").getStatus()); + assertEquals( + ExecutionStatus.SUCCEEDED, runner.runUntilComplete("input").getStatus()); + assertTrue(decisions.size() >= 2, "Initial invocation and actual replay must both execute the handler"); + assertTrue(decisions.stream().allMatch(decision -> decision == expectedSampled), decisions.toString()); + assertEquals( + upstream.equals("root") ? (topology.equals("opaque") ? 1 : decisions.size()) : 0, + evaluations.get()); + assertEquals(1, sideEffects.get()); + } finally { + GlobalOpenTelemetry.resetForTest(); + OtelPluginAutoConfigurationState.resetInstalledForTest(); + DurableSamplingDecision.clearSharedStateForTest(); + if (oldHeader == null) System.clearProperty("com.amazonaws.xray.traceHeader"); + else System.setProperty("com.amazonaws.xray.traceHeader", oldHeader); + } + } +} diff --git a/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/SamplingProcessorIsolationTest.java b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/SamplingProcessorIsolationTest.java new file mode 100644 index 000000000..79981a9b4 --- /dev/null +++ b/otel-plugin/src/test/java/software/amazon/lambda/durable/otel/SamplingProcessorIsolationTest.java @@ -0,0 +1,224 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.otel; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertSame; + +import io.opentelemetry.api.GlobalOpenTelemetry; +import io.opentelemetry.api.OpenTelemetry; +import io.opentelemetry.api.trace.Span; +import io.opentelemetry.api.trace.Tracer; +import io.opentelemetry.api.trace.TracerProvider; +import io.opentelemetry.context.Context; +import io.opentelemetry.context.propagation.ContextPropagators; +import io.opentelemetry.sdk.trace.IdGenerator; +import io.opentelemetry.sdk.trace.ReadWriteSpan; +import io.opentelemetry.sdk.trace.ReadableSpan; +import io.opentelemetry.sdk.trace.SdkTracerProvider; +import io.opentelemetry.sdk.trace.SpanProcessor; +import io.opentelemetry.sdk.trace.samplers.Sampler; +import java.net.URL; +import java.net.URLClassLoader; +import java.time.Instant; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; + +/** + * Verifies the durable sampling decision crosses the application/Java-agent class-loader boundary. + * + *

    Under the documented ADOT setup the plugin JAR is loaded twice: the application class loader computes and stores + * the decision, and a separate Java-agent extension class loader installs and runs the sampler. Because a + * {@link io.opentelemetry.context.ContextKey} uses reference identity, the two loaders hold distinct keys and the + * context carrier alone cannot bridge them. This test reproduces that topology with two child-first class loaders that + * each load {@code DurableSamplingDecision} separately while sharing the OpenTelemetry API/SDK types with the parent, + * then asserts the thread-scoped system-property bridge carries the decision from one loader to the other. + */ +class SamplingProcessorIsolationTest { + + @AfterEach + void clearBridge() { + DurableSamplingDecision.clearSharedStateForTest(); + } + + @ParameterizedTest + @CsvSource({ + "InvocationOtelPlugin,true,false,Invocation", + "InvocationOtelPlugin,false,false,Invocation", + "InvocationOtelPlugin,false,true,Invocation", + "ExecutionOtelPlugin,true,false,Invocation", + "ExecutionOtelPlugin,false,false,Invocation", + "ExecutionOtelPlugin,false,true,Invocation", + "InvocationOtelPlugin,false,false,PlainInvocation", + "ExecutionOtelPlugin,false,false,PlainInvocation" + }) + void agentProcessorForwardingParentCannotReuseApplicationSamplingIntent( + String pluginName, boolean hideProvider, boolean localSampler, String observedSpan) throws Exception { + var plainSampler = observedSpan.equals("PlainInvocation"); + var spanName = plainSampler ? "Invocation" : observedSpan; + var previousHeader = System.getProperty("com.amazonaws.xray.traceHeader"); + GlobalOpenTelemetry.resetForTest(); + OtelPluginAutoConfigurationState.markInstalled(); + System.setProperty("com.amazonaws.xray.traceHeader", "Root=1-6955b900-123456789012345678901234;Sampled=1"); + try (var appLoader = pluginClassLoader(); + var agentLoader = pluginClassLoader()) { + var appSamplerType = Class.forName(DurableSampler.class.getName(), true, appLoader); + var appWrap = appSamplerType.getDeclaredMethod("wrap", Sampler.class); + appWrap.setAccessible(true); + var agentSamplerType = Class.forName(DurableSampler.class.getName(), true, agentLoader); + var agentWrap = agentSamplerType.getDeclaredMethod("wrap", Sampler.class); + agentWrap.setAccessible(true); + var appDecision = Class.forName(DurableSamplingDecision.class.getName(), true, appLoader); + var appGet = appDecision.getDeclaredMethod("get", Context.class); + appGet.setAccessible(true); + var agentDecision = Class.forName(DurableSamplingDecision.class.getName(), true, agentLoader); + var agentGet = agentDecision.getDeclaredMethod("get", Context.class); + agentGet.setAccessible(true); + var agentIdType = Class.forName(DeterministicIdGenerator.class.getName(), true, agentLoader); + var callbacks = new AtomicInteger(); + var leakedSampling = new AtomicBoolean(); + try (var appProvider = SdkTracerProvider.builder() + .setSampler((Sampler) appWrap.invoke(null, Sampler.alwaysOff())) + .build(); + var agentProvider = SdkTracerProvider.builder() + .setSampler( + plainSampler + ? Sampler.alwaysOn() + : (Sampler) (localSampler ? appWrap : agentWrap) + .invoke(null, Sampler.alwaysOff())) + .setIdGenerator( + (IdGenerator) agentIdType.getConstructor().newInstance()) + .addSpanProcessor(new SpanProcessor() { + @Override + public void onStart(Context parent, ReadWriteSpan span) { + if (!span.getName().equals(spanName)) return; + var unrelated = appProvider + .get("processor") + .spanBuilder("unrelated-forwarded-parent") + .setParent(parent) + .startSpan(); + callbacks.incrementAndGet(); + leakedSampling.compareAndSet(false, unrelated.isRecording()); + unrelated.end(); + } + + @Override + public boolean isStartRequired() { + return true; + } + + @Override + public void onEnd(ReadableSpan span) {} + + @Override + public boolean isEndRequired() { + return false; + } + }) + .build()) { + assertEquals(localSampler, agentProvider.getSampler().getClass() == appSamplerType); + var hiddenAgentProvider = new TracerProvider() { + @Override + public Tracer get(String name) { + return agentProvider.get(name); + } + + @Override + public Tracer get(String name, String version) { + return agentProvider.get(name, version); + } + }; + GlobalOpenTelemetry.set(new OpenTelemetry() { + @Override + public TracerProvider getTracerProvider() { + return hideProvider ? hiddenAgentProvider : agentProvider; + } + + @Override + public ContextPropagators getPropagators() { + return ContextPropagators.noop(); + } + }); + var plugin = (DurableExecutionPlugin) + Class.forName("software.amazon.lambda.durable.otel." + pluginName, true, appLoader) + .getConstructor() + .newInstance(); + var arn = "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/test/id"; + var previousAmbient = Span.current(); + var ambient = appProvider + .get("ambient") + .spanBuilder("ambient-control") + .startSpan(); + assertFalse(ambient.isRecording(), "The app provider drops unrelated spans outside the callback"); + try (var ambientScope = ambient.makeCurrent()) { + for (var first : new boolean[] {true, false}) { + plugin.onInvocationStart(new InvocationInfo("request", arn, first, Instant.EPOCH)); + plugin.onInvocationEnd( + new InvocationEndInfo("request", arn, first, InvocationStatus.PENDING, null)); + assertSame(ambient, Span.current(), "Plugin cleanup must preserve the caller's ambient span"); + assertNull(appGet.invoke(null, Context.root()), "Application sampling intent must be cleared"); + assertNull(agentGet.invoke(null, Context.root()), "Agent sampling intent must be cleared"); + } + } finally { + ambient.end(); + } + assertSame(previousAmbient, Span.current()); + assertEquals(2, callbacks.get(), "The selected span in both invocations must reach the real processor"); + assertFalse( + leakedSampling.get(), "The supplied parent must not override the app provider's DROP policy"); + } + } finally { + GlobalOpenTelemetry.resetForTest(); + OtelPluginAutoConfigurationState.resetInstalledForTest(); + if (previousHeader == null) System.clearProperty("com.amazonaws.xray.traceHeader"); + else System.setProperty("com.amazonaws.xray.traceHeader", previousHeader); + } + } + + /** + * A child-first class loader that loads {@code software.amazon.lambda.durable.otel.*} itself (so each instance + * holds its own copies, mirroring the two plugin class loaders) while delegating OpenTelemetry and JDK classes to + * the parent so those types are shared and interoperable across loaders. + */ + private static URLClassLoader pluginClassLoader() { + var classesDir = SamplingProcessorIsolationTest.class + .getProtectionDomain() + .getCodeSource() + .getLocation(); + // target/test-classes -> the main classes live in target/classes alongside it. + URL mainClasses; + try { + mainClasses = new URL(classesDir.toString().replace("/test-classes/", "/classes/")); + } catch (Exception e) { + throw new IllegalStateException(e); + } + var parent = SamplingProcessorIsolationTest.class.getClassLoader(); + return new URLClassLoader(new URL[] {mainClasses}, parent) { + @Override + protected Class loadClass(String name, boolean resolve) throws ClassNotFoundException { + if (name.startsWith("software.amazon.lambda.durable.otel.")) { + synchronized (getClassLoadingLock(name)) { + var loaded = findLoadedClass(name); + if (loaded == null) { + loaded = findClass(name); + } + if (resolve) { + resolveClass(loaded); + } + return loaded; + } + } + return super.loadClass(name, resolve); + } + }; + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/ContinuationFailurePublicProbeTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/ContinuationFailurePublicProbeTest.java new file mode 100644 index 000000000..7837527b5 --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/ContinuationFailurePublicProbeTest.java @@ -0,0 +1,241 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Set; +import java.util.concurrent.*; +import java.util.concurrent.atomic.*; +import java.util.function.BiFunction; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.config.WaitForConditionConfig; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.execution.ExecutionManager; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.DurableExecutionOutput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.model.WaitForConditionResult; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class ContinuationFailurePublicProbeTest { + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void readyResumeFailureDoesNotDisappearOrHang(boolean rejection, boolean withPlugin) throws Exception { + RuntimeException failure = rejection + ? new RejectedExecutionException("resumed worker rejected") + : new IllegalArgumentException("resumed state cannot deserialize"); + var managerForCleanup = new AtomicReference(); + var armed = new AtomicBoolean(); + var injected = new CountDownLatch(1); + var injectionOwner = new AtomicReference(); + var checks = new AtomicInteger(); + var ends = new CopyOnWriteArrayList(); + var pollEntered = new CountDownLatch(1); + var releaseReady = new CountDownLatch(1); + var firstArmedPoll = new AtomicBoolean(true); + var onCoordinator = new AtomicBoolean(); + var pendingCall = new AtomicReference>(); + var serde = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + public String serialize(Object value) { + return delegate.serialize(value); + } + + public T deserialize(String value, TypeToken type) { + if (!rejection + && armed.get() + && Thread.currentThread().getName().startsWith("durable-sdk-internal-")) { + injectionOwner.set(Thread.currentThread().getName()); + onCoordinator.set(Arrays.stream(Thread.currentThread().getStackTrace()) + .anyMatch(frame -> frame.getMethodName().equals("completeCheckpointContinuation"))); + injected.countDown(); + throw failure; + } + return delegate.deserialize(value, type); + } + }; + var client = new LocalMemoryExecutionClient() { + @Override + public CheckpointDurableExecutionResponse checkpoint( + String arn, String token, List updates) { + if (armed.get() && updates.isEmpty() && firstArmedPoll.compareAndSet(true, false)) { + pollEntered.countDown(); + try { + assertTrue(releaseReady.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + throw new AssertionError(interrupted); + } + } + return super.checkpoint(arn, token, updates); + } + }; + var workers = + new ThreadPoolExecutor(0, Integer.MAX_VALUE, 60, TimeUnit.SECONDS, new SynchronousQueue()) { + @Override + public void execute(Runnable task) { + if (rejection + && armed.get() + && Thread.currentThread().getName().startsWith("durable-sdk-internal-")) { + injectionOwner.set(Thread.currentThread().getName()); + onCoordinator.set(Arrays.stream( + Thread.currentThread().getStackTrace()) + .anyMatch(frame -> frame.getMethodName().equals("completeCheckpointContinuation"))); + injected.countDown(); + throw failure; + } + super.execute(task); + } + }; + var caller = Executors.newSingleThreadExecutor(); + var builder = DurableConfig.builder() + .withExecutorService(workers) + .withDurableExecutionClient(client) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO); + if (withPlugin) + builder.withPlugins(new DurableExecutionPlugin() { + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + } + }); + var config = builder.build(); + var condition = WaitForConditionConfig.builder() + .initialState(1) + .serDes(serde) + .waitStrategy((state, attempt) -> Duration.ofSeconds(1)) + .build(); + BiFunction handler = (input, context) -> { + managerForCleanup.set( + ((software.amazon.lambda.durable.context.BaseContextImpl) context).getExecutionManager()); + var waiting = context.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + checks.incrementAndGet(); + return state == 2 + ? WaitForConditionResult.stopPolling(state) + : WaitForConditionResult.continuePolling(2); + }, + condition); + if (armed.get()) { + try { + assertTrue(pollEntered.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + throw new AssertionError(interrupted); + } + } + return String.valueOf(waiting.get()); + }; + try { + var first = caller.submit(() -> DurableExecutor.execute( + input(List.of()), null, TypeToken.get(String.class), handler, config)) + .get(3, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.PENDING, first.status()); + var cachedPending = new ArrayList<>(client.getAllOperations()); + assertEquals( + OperationStatus.PENDING, + client.getOperationByName("condition").status()); + assertTrue(client.advanceTime()); + armed.set(true); + var second = caller.submit(() -> + DurableExecutor.execute(input(cachedPending), null, TypeToken.get(String.class), handler, config)); + pendingCall.set(second); + assertTrue(pollEntered.await(3, TimeUnit.SECONDS)); + releaseReady.countDown(); + assertTrue(injected.await(3, TimeUnit.SECONDS)); + assertTrue(onCoordinator.get(), "Fault must run in the actual WFC coordinator continuation"); + DurableExecutionOutput output = null; + Throwable callerFailure = null; + try { + output = second.get(3, TimeUnit.SECONDS); + } catch (ExecutionException observationFailure) { + callerFailure = observationFailure.getCause(); + } + System.out.println("CONTINUATION_ORDINARY_PUBLIC owner=" + injectionOwner.get() + + " caller=" + + (output != null + ? output.status() + : callerFailure.getClass().getName()) + + " backend=" + client.getOperationByName("condition").status() + + " checks=" + checks.get()); + assertEquals(1, checks.get(), "A failed resumed state read must not rerun the condition body"); + assertNull(output, "An SDK continuation failure must not produce a durable response"); + var retry = assertInstanceOf(UnrecoverableDurableExecutionException.class, callerFailure); + assertTrue(retry.isRetryable()); + assertSame(failure, retry.getCause()); + assertEquals( + OperationStatus.READY, + client.getOperationByName("condition").status()); + var activeField = ExecutionManager.class.getDeclaredField("activeThreads"); + activeField.setAccessible(true); + assertFalse( + ((Set) activeField.get(managerForCleanup.get())) + .contains(client.getOperationByName("condition").id()), + "A rejected operation worker must not leave a phantom activity registration"); + if (withPlugin) + assertEquals( + List.of(InvocationStatus.PENDING, InvocationStatus.RETRYING), + ends.stream().map(InvocationEndInfo::invocationStatus).toList()); + else assertTrue(ends.isEmpty()); + armed.set(false); + var third = caller.submit(() -> DurableExecutor.execute( + input(client.getAllOperations()), null, TypeToken.get(String.class), handler, config)); + pendingCall.set(third); + assertEquals( + ExecutionStatus.SUCCEEDED, third.get(3, TimeUnit.SECONDS).status()); + assertEquals(2, checks.get(), "Only the original and successful resumed predicates run"); + } finally { + releaseReady.countDown(); + // Cleanup-only escape hatch AFTER the real public invocation observation/timeout; never test behavior. + var manager = managerForCleanup.get(); + if (pendingCall.get() != null + && !pendingCall.get().isDone() + && manager != null + && !manager.isExecutionCompletedExceptionally()) { + try { + manager.terminateExecution(new UnrecoverableDurableExecutionException( + ErrorObject.builder().errorMessage("probe cleanup").build(), true)); + } catch (UnrecoverableDurableExecutionException expected) { + } + } + caller.shutdownNow(); + workers.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + private static DurableExecutionInput input(List operations) { + var execution = Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + var history = new ArrayList(); + history.add(execution); + history.addAll(operations); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/fatal/execution", + "token", + CheckpointUpdatedExecutionState.builder().operations(history).build()); + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/ContinuationFatalPublicProbeTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/ContinuationFatalPublicProbeTest.java new file mode 100644 index 000000000..43e76b8cb --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/ContinuationFatalPublicProbeTest.java @@ -0,0 +1,233 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.*; +import java.util.concurrent.atomic.*; +import java.util.function.BiFunction; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.config.WaitForConditionConfig; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.DurableExecutionOutput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.model.WaitForConditionResult; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class ContinuationFatalPublicProbeTest { + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void readyResumeSerdeFatalDoesNotDisappearIntoPending(boolean death, boolean withPlugin) throws Exception { + runFatalWithEndFailure(death, withPlugin, null); + } + + @ParameterizedTest + @CsvSource({"false,assertion", "true,assertion", "false,fatal", "true,fatal"}) + void continuationFatalKeepsDiagnosticsWhenEndAlsoFails(boolean death, String endKind) throws Exception { + runFatalWithEndFailure(death, true, endKind); + } + + @SuppressWarnings("removal") + private static void runFatalWithEndFailure(boolean death, boolean withPlugin, String endKind) throws Exception { + Error endFailure = endKind == null + ? null + : endKind.equals("fatal") + ? (death ? new InternalError("End fatal") : new ThreadDeath()) + : new AssertionError("End assertion"); + var endEscaped = new CountDownLatch(1); + var endOwnerFailure = new AtomicReference(); + Error fatal = death ? new ThreadDeath() : new InternalError("resume serde fatal"); + var armed = new AtomicBoolean(); + var injected = new CountDownLatch(1); + var escaped = new CountDownLatch(1); + var injectionOwner = new AtomicReference(); + var ownerFailure = new AtomicReference(); + var checks = new AtomicInteger(); + var ends = new CopyOnWriteArrayList(); + var pollEntered = new CountDownLatch(1); + var releaseReady = new CountDownLatch(1); + var firstArmedPoll = new AtomicBoolean(true); + var onCoordinator = new AtomicBoolean(); + var previousHandler = Thread.getDefaultUncaughtExceptionHandler(); + Thread.setDefaultUncaughtExceptionHandler((thread, failure) -> { + if (failure == fatal) { + ownerFailure.set(failure); + escaped.countDown(); + } else if (failure == endFailure) { + endOwnerFailure.set(failure); + endEscaped.countDown(); + } else if (previousHandler != null) previousHandler.uncaughtException(thread, failure); + }); + var serde = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + public String serialize(Object value) { + return delegate.serialize(value); + } + + public T deserialize(String value, TypeToken type) { + if (armed.get() && Thread.currentThread().getName().startsWith("durable-sdk-internal-")) { + injectionOwner.set(Thread.currentThread().getName()); + onCoordinator.set(Arrays.stream(Thread.currentThread().getStackTrace()) + .anyMatch(frame -> frame.getMethodName().equals("completeCheckpointContinuation"))); + injected.countDown(); + throw fatal; + } + return delegate.deserialize(value, type); + } + }; + var client = new LocalMemoryExecutionClient() { + @Override + public CheckpointDurableExecutionResponse checkpoint( + String arn, String token, List updates) { + if (armed.get() && updates.isEmpty() && firstArmedPoll.compareAndSet(true, false)) { + pollEntered.countDown(); + try { + assertTrue(releaseReady.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException failure) { + throw new AssertionError(failure); + } + } + return super.checkpoint(arn, token, updates); + } + }; + var workers = Executors.newCachedThreadPool(); + var caller = Executors.newSingleThreadExecutor(); + var builder = DurableConfig.builder() + .withExecutorService(workers) + .withDurableExecutionClient(client) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO); + if (withPlugin) + builder.withPlugins(new DurableExecutionPlugin() { + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + if (info.invocationStatus() == InvocationStatus.RETRYING && endFailure != null) throw endFailure; + } + }); + var config = builder.build(); + var condition = WaitForConditionConfig.builder() + .initialState(1) + .serDes(serde) + .waitStrategy((state, attempt) -> Duration.ofSeconds(1)) + .build(); + BiFunction handler = (input, context) -> { + var waiting = context.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + checks.incrementAndGet(); + return state == 2 + ? WaitForConditionResult.stopPolling(state) + : WaitForConditionResult.continuePolling(2); + }, + condition); + if (armed.get()) { + try { + assertTrue(pollEntered.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException failure) { + throw new AssertionError(failure); + } + } + return String.valueOf(waiting.get()); + }; + try { + var first = caller.submit(() -> DurableExecutor.execute( + input(List.of()), null, TypeToken.get(String.class), handler, config)) + .get(3, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.PENDING, first.status()); + var cachedPending = new ArrayList<>(client.getAllOperations()); + assertEquals( + OperationStatus.PENDING, + client.getOperationByName("condition").status()); + assertTrue(client.advanceTime()); + armed.set(true); + var second = caller.submit(() -> + DurableExecutor.execute(input(cachedPending), null, TypeToken.get(String.class), handler, config)); + assertTrue(pollEntered.await(3, TimeUnit.SECONDS)); + releaseReady.countDown(); + assertTrue(injected.await(3, TimeUnit.SECONDS)); + assertTrue(onCoordinator.get(), "Fault must run in the actual WFC coordinator continuation"); + DurableExecutionOutput output = null; + Throwable callerFailure = null; + try { + output = second.get(3, TimeUnit.SECONDS); + } catch (ExecutionException failure) { + callerFailure = failure.getCause(); + } + var ownerEscaped = escaped.await(300, TimeUnit.MILLISECONDS); + System.out.println("CONTINUATION_FATAL_PUBLIC owner=" + injectionOwner.get() + + " caller=" + + (output != null + ? output.status() + : callerFailure.getClass().getName()) + + " backend=" + client.getOperationByName("condition").status() + + " checks=" + checks.get() + " workerEscaped=" + ownerEscaped); + assertEquals(1, checks.get(), "A failed resumed state read must not rerun the condition body"); + assertTrue(ownerEscaped, "The actual coordinator worker must observe its JVM fatal"); + assertSame(fatal, ownerFailure.get()); + assertNull(output, "A failed coordinator must not return a durable terminal or PENDING response"); + UnrecoverableDurableExecutionException retry; + if ("fatal".equals(endKind)) { + assertSame(endFailure, callerFailure, "Existing Root End JVM-fatal precedence remains"); + assertTrue(endEscaped.await(3, TimeUnit.SECONDS)); + assertSame(endFailure, endOwnerFailure.get()); + assertEquals(1, endFailure.getSuppressed().length); + retry = assertInstanceOf( + UnrecoverableDurableExecutionException.class, endFailure.getSuppressed()[0]); + } else { + retry = assertInstanceOf(UnrecoverableDurableExecutionException.class, callerFailure); + if (endFailure != null) assertArrayEquals(new Throwable[] {endFailure}, retry.getSuppressed()); + } + assertTrue(retry.isRetryable()); + assertSame(fatal, retry.getCause()); + if (withPlugin) { + assertEquals( + List.of(InvocationStatus.PENDING, InvocationStatus.RETRYING), + ends.stream().map(InvocationEndInfo::invocationStatus).toList()); + assertSame(retry, ends.get(1).executionError()); + } else assertTrue(ends.isEmpty()); + } finally { + releaseReady.countDown(); + Thread.setDefaultUncaughtExceptionHandler(previousHandler); + caller.shutdownNow(); + workers.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + private static DurableExecutionInput input(List operations) { + var execution = Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + var history = new ArrayList(); + history.add(execution); + history.addAll(operations); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/fatal/execution", + "token", + CheckpointUpdatedExecutionState.builder().operations(history).build()); + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/InvocationFinalizationIntegrationTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/InvocationFinalizationIntegrationTest.java new file mode 100644 index 000000000..0e41ada4b --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/InvocationFinalizationIntegrationTest.java @@ -0,0 +1,415 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.time.Duration; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.concurrent.AbstractExecutorService; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.concurrent.locks.LockSupport; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.ValueSource; +import software.amazon.awssdk.services.lambda.model.ErrorObject; +import software.amazon.lambda.durable.config.StepConfig; +import software.amazon.lambda.durable.config.WaitForConditionConfig; +import software.amazon.lambda.durable.context.DurableContextImpl; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.execution.ExecutionManager; +import software.amazon.lambda.durable.execution.SuspendExecutionException; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.model.WaitForConditionResult; +import software.amazon.lambda.durable.operation.BaseDurableOperation; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.retry.RetryStrategies; +import software.amazon.lambda.durable.testing.LocalDurableTestRunner; + +class InvocationFinalizationIntegrationTest { + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void suspensionAndRetryWaitForHandlerFinallyAndOwnerEnd(boolean retry) throws Exception { + var enteredFinally = new CountDownLatch(1); + var releaseFinally = new CountDownLatch(1); + var plugin = new RecordingPlugin(); + var original = retryError(); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + try { + if (retry) + context.step("retry", String.class, step -> { + throw original; + }); + else context.wait("pause", Duration.ofSeconds(1)); + return "done"; + } finally { + enteredFinally.countDown(); + await(releaseFinally); + } + }, + DurableConfig.builder().withPlugins(plugin).build()); + var caller = Executors.newSingleThreadExecutor(); + try { + var response = caller.submit(() -> runner.run("input")); + assertTrue(enteredFinally.await(5, TimeUnit.SECONDS)); + assertThrows(TimeoutException.class, () -> response.get(700, TimeUnit.MILLISECONDS)); + assertTrue(plugin.ends.isEmpty(), "End must follow the handler's finally block"); + releaseFinally.countDown(); + if (retry) { + var thrown = assertThrows(ExecutionException.class, () -> response.get(5, TimeUnit.SECONDS)); + assertSame(original, thrown.getCause()); + } else + assertEquals( + ExecutionStatus.PENDING, + response.get(5, TimeUnit.SECONDS).getStatus()); + assertPaired(plugin, retry ? InvocationStatus.RETRYING : InvocationStatus.PENDING); + } finally { + releaseFinally.countDown(); + stop(caller); + } + } + + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void externallySelectedOutcomeSurvivesReturningOrThrowingFinally(boolean retry, boolean throwsInFinally) + throws Exception { + var enteredFinally = new CountDownLatch(1); + var releaseFinally = new CountDownLatch(1); + var manager = new AtomicReference(); + var plugin = new RecordingPlugin(); + var original = retryError(); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + manager.set(((DurableContextImpl) context).getExecutionManager()); + try { + return "handler completed"; + } finally { + enteredFinally.countDown(); + await(releaseFinally); + if (throwsInFinally) throw new IllegalStateException("finally failed"); + } + }, + DurableConfig.builder().withPlugins(plugin).build()); + var caller = Executors.newSingleThreadExecutor(); + try { + var response = caller.submit(() -> runner.run("input")); + assertTrue(enteredFinally.await(5, TimeUnit.SECONDS)); + if (retry) + assertSame( + original, + assertThrows( + UnrecoverableDurableExecutionException.class, + () -> manager.get().terminateExecution(original))); + else + assertThrows( + SuspendExecutionException.class, () -> manager.get().suspendExecution()); + assertThrows(TimeoutException.class, () -> response.get(700, TimeUnit.MILLISECONDS)); + releaseFinally.countDown(); + if (retry) { + var thrown = assertThrows(ExecutionException.class, () -> response.get(5, TimeUnit.SECONDS)); + assertSame(original, thrown.getCause()); + } else + assertEquals( + ExecutionStatus.PENDING, + response.get(5, TimeUnit.SECONDS).getStatus()); + assertPaired(plugin, retry ? InvocationStatus.RETRYING : InvocationStatus.PENDING); + assertSame(retry ? original : null, plugin.ends.get(0).executionError()); + assertNull(plugin.ends.get(0).executionResult()); + } finally { + releaseFinally.countDown(); + stop(caller); + } + } + + @ParameterizedTest + @CsvSource({ + "step,async,true", + "step,async,false", + "condition,async,true", + "condition,async,false", + "termination,async,true", + "termination,async,false", + "step,inline-root,true", + "step,inline-root,false", + "condition,inline-root,true", + "condition,inline-root,false", + "termination,inline-root,true", + "termination,inline-root,false", + "step,submit-wait,true", + "step,submit-wait,false", + "condition,submit-wait,true", + "condition,submit-wait,false", + "termination,submit-wait,true", + "termination,submit-wait,false" + }) + void naturalManagerOutcomePrecedesFastThrowingFinally(String kind, String mode, boolean withPlugin) + throws Exception { + var workers = new OutcomeExecutor(mode); + var releaseOperation = new CountDownLatch(mode.equals("submit-wait") ? 0 : 1); + var aboutToWait = new CountDownLatch(1); + var completedStepCalls = new AtomicInteger(); + var targetCalls = new AtomicInteger(); + var throwFinally = new AtomicBoolean(true); + var original = retryError(); + var plugin = new RecordingPlugin(); + var builder = DurableConfig.builder().withExecutorService(workers); + if (withPlugin) builder.withPlugins(plugin); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + context.step("saved", String.class, step -> { + completedStepCalls.incrementAndGet(); + return "saved"; + }); + DurableFuture future; + if (kind.equals("condition")) { + future = context.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + await(releaseOperation); + return targetCalls.incrementAndGet() == 1 + ? WaitForConditionResult.continuePolling(state + 1) + : WaitForConditionResult.stopPolling(state); + }, + WaitForConditionConfig.builder() + .initialState(0) + .waitStrategy((state, attempt) -> Duration.ofSeconds(1)) + .build()); + } else { + future = context.stepAsync( + "target", + String.class, + step -> { + await(releaseOperation); + if (targetCalls.incrementAndGet() == 1) { + if (kind.equals("termination")) throw original; + throw new IllegalArgumentException("retry step"); + } + return "target"; + }, + StepConfig.builder() + .retryStrategy(RetryStrategies.fixedDelay(2, Duration.ofSeconds(1))) + .build()); + } + // This schedules the existing completion window; suspension/termination is triggered only by the + // real + // operation. The root's dependent wait is registered later and wakes before this earlier callback + // drains. + ((BaseDurableOperation) future).getCompletionFuture().whenComplete((value, failure) -> { + if (failure != null && !mode.equals("submit-wait") && throwFinally.get()) + await(workers.rootReturned); + }); + try { + aboutToWait.countDown(); + future.get(); + return "done"; + } finally { + if (throwFinally.get()) throw new IllegalStateException("fast handler finally"); + } + }, + builder.build()); + var caller = Executors.newSingleThreadExecutor(); + try { + var response = caller.submit(() -> runner.run("input")); + await(aboutToWait); + if (!mode.equals("submit-wait")) { + var until = System.nanoTime() + TimeUnit.SECONDS.toNanos(3); + while (System.nanoTime() < until) { + var root = workers.rootOwner.get(); + if (root != null + && root.getState() == Thread.State.WAITING + && Arrays.stream(root.getStackTrace()) + .anyMatch(frame -> frame.getMethodName().equals("waitForOperationCompletion"))) + break; + LockSupport.parkNanos(100_000); + } + assertEquals(Thread.State.WAITING, workers.rootOwner.get().getState()); + releaseOperation.countDown(); + } + if (kind.equals("termination")) { + assertSame( + original, + assertThrows(ExecutionException.class, () -> response.get(5, TimeUnit.SECONDS)) + .getCause()); + } else { + assertEquals( + ExecutionStatus.PENDING, + response.get(5, TimeUnit.SECONDS).getStatus()); + } + if (withPlugin) + assertPaired(plugin, kind.equals("termination") ? InvocationStatus.RETRYING : InvocationStatus.PENDING); + throwFinally.set(false); + assertEquals( + ExecutionStatus.SUCCEEDED, runner.runUntilComplete("input").getStatus()); + assertEquals(1, completedStepCalls.get(), "The completed step must replay without another side effect"); + } finally { + releaseOperation.countDown(); + workers.rootReturned.countDown(); + stop(caller); + stop(workers); + } + } + + private static final class OutcomeExecutor extends AbstractExecutorService { + private final String mode; + private final ExecutorService delegate = Executors.newCachedThreadPool(); + private final AtomicBoolean first = new AtomicBoolean(true); + private final AtomicReference rootOwner = new AtomicReference<>(); + private final CountDownLatch rootReturned = new CountDownLatch(1); + + private OutcomeExecutor(String mode) { + this.mode = mode; + } + + public void execute(Runnable task) { + var root = first.getAndSet(false); + Runnable wrapped = () -> { + if (root) rootOwner.set(Thread.currentThread()); + try { + task.run(); + } finally { + if (root) rootReturned.countDown(); + } + }; + if (root && mode.equals("inline-root")) wrapped.run(); + else if (mode.equals("submit-wait")) { + try { + delegate.submit(wrapped).get(); + } catch (Exception failure) { + throw new AssertionError(failure); + } + } else delegate.execute(wrapped); + } + + public void shutdown() { + delegate.shutdown(); + } + + public List shutdownNow() { + return delegate.shutdownNow(); + } + + public boolean isShutdown() { + return delegate.isShutdown(); + } + + public boolean isTerminated() { + return delegate.isTerminated(); + } + + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(timeout, unit); + } + } + + @Test + void hooksStayPairedAcrossStepRetryAndReplayWithoutRepeatingCompletedWork() { + var plugin = new RecordingPlugin(); + var sideEffects = new AtomicInteger(); + var attempts = new AtomicInteger(); + var runner = LocalDurableTestRunner.create( + String.class, + (input, context) -> { + var result = context.step("completed", Integer.class, step -> sideEffects.incrementAndGet()); + context.step( + "retry", + String.class, + step -> { + if (attempts.incrementAndGet() == 1) throw new IllegalStateException("try again"); + return "retried"; + }, + StepConfig.builder() + .retryStrategy(RetryStrategies.fixedDelay(2, Duration.ofSeconds(1))) + .build()); + context.wait("pause", Duration.ofSeconds(1)); + return result; + }, + DurableConfig.builder().withPlugins(plugin).build()); + + var output = runner.runUntilComplete("input"); + + assertEquals(ExecutionStatus.SUCCEEDED, output.getStatus()); + assertEquals(1, output.getResult(Integer.class)); + assertEquals(1, sideEffects.get(), "the completed step must replay without repeating user code"); + assertEquals(2, attempts.get()); + assertTrue(plugin.starts.size() >= 3, "retry and wait should each suspend before the successful invocation"); + assertEquals(plugin.starts.size(), plugin.ends.size()); + assertEquals(plugin.startThreads, plugin.endThreads); + assertTrue(plugin.endLocalValues.stream().allMatch("invocation"::equals)); + assertEquals(InvocationStatus.PENDING, plugin.ends.get(0).invocationStatus()); + assertEquals( + InvocationStatus.SUCCEEDED, + plugin.ends.get(plugin.ends.size() - 1).invocationStatus()); + assertTrue(plugin.starts.get(0).isFirstInvocation()); + assertFalse(plugin.starts.get(plugin.starts.size() - 1).isFirstInvocation()); + } + + private static void assertPaired(RecordingPlugin plugin, InvocationStatus status) { + assertEquals(1, plugin.starts.size()); + assertEquals(1, plugin.ends.size()); + assertEquals(status, plugin.ends.get(0).invocationStatus()); + assertEquals(plugin.startThreads, plugin.endThreads); + assertEquals(List.of("invocation"), plugin.endLocalValues); + } + + private static UnrecoverableDurableExecutionException retryError() { + return new UnrecoverableDurableExecutionException( + ErrorObject.builder().errorMessage("retry invocation").build(), true); + } + + private static void await(CountDownLatch latch) { + try { + if (!latch.await(5, TimeUnit.SECONDS)) throw new AssertionError("finally not released"); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + + private static void stop(ExecutorService executor) throws InterruptedException { + executor.shutdownNow(); + assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS)); + } + + private static class RecordingPlugin implements DurableExecutionPlugin { + final ThreadLocal local = new ThreadLocal<>(); + final List starts = new ArrayList<>(); + final List ends = new ArrayList<>(); + final List startThreads = new ArrayList<>(); + final List endThreads = new ArrayList<>(); + final List endLocalValues = new ArrayList<>(); + + @Override + public void onInvocationStart(InvocationInfo info) { + starts.add(info); + startThreads.add(Thread.currentThread()); + local.set("invocation"); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + endThreads.add(Thread.currentThread()); + endLocalValues.add(local.get()); + local.remove(); + } + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/PreparationFailureDiagnosticsTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/PreparationFailureDiagnosticsTest.java new file mode 100644 index 000000000..66528416d --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/PreparationFailureDiagnosticsTest.java @@ -0,0 +1,111 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.time.Duration; +import java.time.Instant; +import java.util.List; +import java.util.concurrent.*; +import java.util.concurrent.atomic.*; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.plugin.*; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class PreparationFailureDiagnosticsTest { + @ParameterizedTest + @ValueSource(strings = {"serialize", "checkpoint"}) + void preparationCauseDeliveredToCallerRetainsEndError(String mode) throws Exception { + var preparation = new IllegalStateException("output preparation failed"); + var endFailure = new AssertionError("End error"); + var end = new AtomicReference(); + var endCalls = new AtomicInteger(); + var checkpoints = new AtomicInteger(); + var workers = Executors.newCachedThreadPool(); + var caller = Executors.newSingleThreadExecutor(); + var client = new LocalMemoryExecutionClient() { + @Override + public CheckpointDurableExecutionResponse checkpoint( + String arn, String token, List updates) { + if (mode.equals("checkpoint") + && updates.stream().anyMatch(op -> op.type() == OperationType.EXECUTION)) { + checkpoints.incrementAndGet(); + throw preparation; + } + return super.checkpoint(arn, token, updates); + } + }; + var serDes = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + public String serialize(Object value) { + if (mode.equals("serialize") && "result".equals(value)) throw preparation; + return delegate.serialize(value); + } + + public T deserialize(String value, TypeToken type) { + return delegate.deserialize(value, type); + } + }; + var config = DurableConfig.builder() + .withDurableExecutionClient(client) + .withExecutorService(workers) + .withSerDes(serDes) + .withCheckpointDelay(Duration.ZERO) + .withPlugins(new DurableExecutionPlugin() { + public void onInvocationEnd(InvocationEndInfo info) { + endCalls.incrementAndGet(); + end.set(info); + throw endFailure; + } + }) + .build(); + try { + var result = caller.submit(() -> DurableExecutor.execute( + input(), + null, + TypeToken.get(String.class), + (value, context) -> mode.equals("checkpoint") ? "x".repeat(6 * 1024 * 1024) : "result", + config)); + var observed = assertThrows(ExecutionException.class, () -> result.get(5, TimeUnit.SECONDS)) + .getCause(); + System.out.println("PREPARATION_DIAGNOSTICS mode=" + mode + " original=" + (observed == preparation) + + " suppressed=" + List.of(observed.getSuppressed()) + " End=" + endCalls.get()); + assertSame(preparation, observed); + assertEquals(List.of(endFailure), List.of(observed.getSuppressed())); + assertEquals(1, endCalls.get()); + assertEquals(InvocationStatus.RETRYING, end.get().invocationStatus()); + assertSame(preparation, end.get().executionError()); + assertEquals(mode.equals("checkpoint") ? 1 : 0, checkpoints.get()); + } finally { + caller.shutdownNow(); + workers.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + private static DurableExecutionInput input() { + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/diagnostics/execution", + "token", + CheckpointUpdatedExecutionState.builder() + .operations(Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails(ExecutionDetails.builder() + .inputPayload("\"input\"") + .build()) + .build()) + .build()); + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/StepRetryContinuationTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/StepRetryContinuationTest.java new file mode 100644 index 000000000..ec723205c --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/StepRetryContinuationTest.java @@ -0,0 +1,483 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.lang.management.ManagementFactory; +import java.lang.reflect.Method; +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Deque; +import java.util.List; +import java.util.Map; +import java.util.concurrent.*; +import java.util.concurrent.atomic.*; +import java.util.concurrent.locks.LockSupport; +import java.util.function.BiFunction; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.ValueSource; +import org.slf4j.MDC; +import org.slf4j.spi.MDCAdapter; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.config.StepConfig; +import software.amazon.lambda.durable.context.BaseContextImpl; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.execution.ExecutionManager; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.DurableExecutionOutput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.retry.RetryDecision; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class StepRetryContinuationTest { + @ParameterizedTest + @ValueSource(strings = {"async", "reject", "direct", "submit-wait"}) + void realPendingStepRetryMustNotLoseDispatchFailureOrBlockItsCheckpointBatcher(String mode) throws Exception { + var armed = new AtomicBoolean(); + var calls = new AtomicInteger(); + var pollEntered = new CountDownLatch(1); + var releasePoll = new CountDownLatch(1); + var pollOnce = new AtomicBoolean(true); + var dispatchTimedOut = new AtomicBoolean(); + var dispatchThread = new AtomicReference(); + var dispatchFailure = new RejectedExecutionException("resumed step rejected"); + var managerRef = new AtomicReference(); + var pendingCall = new AtomicReference>(); + var client = new LocalMemoryExecutionClient() { + @Override + public CheckpointDurableExecutionResponse checkpoint( + String arn, String token, List updates) { + if (armed.get() && updates.isEmpty() && pollOnce.compareAndSet(true, false)) { + pollEntered.countDown(); + await(releasePoll); + } + return super.checkpoint(arn, token, updates); + } + }; + var workers = + new ThreadPoolExecutor(0, Integer.MAX_VALUE, 60, TimeUnit.SECONDS, new SynchronousQueue()) { + @Override + public void execute(Runnable task) { + boolean resume = armed.get() + && Arrays.stream(Thread.currentThread().getStackTrace()) + .anyMatch(frame -> frame.getClassName().endsWith(".StepOperation") + && frame.getMethodName().equals("executeStepLogic")); + if (resume) { + dispatchThread.set(Thread.currentThread().getName()); + if (mode.equals("reject")) throw dispatchFailure; + if (mode.equals("direct")) { + task.run(); + return; + } + if (mode.equals("submit-wait")) { + var done = new CompletableFuture(); + super.execute(() -> { + try { + task.run(); + } finally { + done.complete(null); + } + }); + try { + done.get(2, TimeUnit.SECONDS); + } catch (TimeoutException expected) { + dispatchTimedOut.set(true); + } catch (Exception failure) { + throw new AssertionError(failure); + } + return; + } + } + super.execute(task); + } + }; + var config = DurableConfig.builder() + .withExecutorService(workers) + .withDurableExecutionClient(client) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO) + .build(); + var stepConfig = StepConfig.builder() + .retryStrategy((error, attempt) -> + attempt < 2 ? RetryDecision.retry(Duration.ofSeconds(1)) : RetryDecision.fail()) + .build(); + var caller = Executors.newSingleThreadExecutor(); + BiFunction handler = (value, context) -> { + managerRef.set(((BaseContextImpl) context).getExecutionManager()); + var future = context.stepAsync( + "retry-step", + String.class, + step -> { + if (calls.incrementAndGet() == 1) throw new IllegalStateException("first attempt"); + return "stored"; + }, + stepConfig); + if (armed.get()) await(pollEntered); + return future.get(); + }; + try { + var first = caller.submit(() -> + DurableExecutor.execute(input(List.of()), null, TypeToken.get(String.class), handler, config)); + pendingCall.set(first); + assertEquals(ExecutionStatus.PENDING, first.get(5, TimeUnit.SECONDS).status()); + var pending = new ArrayList<>(client.getAllOperations()); + assertEquals( + OperationStatus.PENDING, + client.getOperationByName("retry-step").status()); + assertTrue(client.advanceTime()); + armed.set(true); + var second = caller.submit( + () -> DurableExecutor.execute(input(pending), null, TypeToken.get(String.class), handler, config)); + pendingCall.set(second); + assertTrue(pollEntered.await(5, TimeUnit.SECONDS)); + releasePoll.countDown(); + DurableExecutionOutput output = null; + Throwable failure = null; + try { + output = second.get(6, TimeUnit.SECONDS); + } catch (ExecutionException observed) { + failure = observed.getCause(); + } + System.out.println("STEP_RETRY_PUBLIC mode=" + mode + " thread=" + dispatchThread.get() + " output=" + + (output == null ? failure : output.status()) + " backend=" + + client.getOperationByName("retry-step").status() + + " calls=" + calls.get() + " dispatchTimedOut=" + dispatchTimedOut.get()); + assertFalse( + dispatchTimedOut.get(), "Submit-and-wait dispatch must not block the batcher needed by its step"); + if (mode.equals("reject")) { + var retry = assertInstanceOf(UnrecoverableDurableExecutionException.class, failure); + assertTrue(retry.isRetryable()); + assertSame(dispatchFailure, retry.getCause()); + assertNull(output); + assertEquals(1, calls.get()); + assertEquals( + OperationStatus.READY, + client.getOperationByName("retry-step").status()); + } else assertEquals(ExecutionStatus.SUCCEEDED, output.status()); + armed.set(false); + var resumed = caller.submit(() -> DurableExecutor.execute( + input(client.getAllOperations()), null, TypeToken.get(String.class), handler, config)); + pendingCall.set(resumed); + assertEquals( + ExecutionStatus.SUCCEEDED, resumed.get(5, TimeUnit.SECONDS).status()); + assertEquals(2, calls.get(), "Completed step bodies do not repeat on replay"); + } finally { + releasePoll.countDown(); + var manager = managerRef.get(); + if (pendingCall.get() != null + && !pendingCall.get().isDone() + && manager != null + && !manager.isExecutionCompletedExceptionally()) { + try { + manager.terminateExecution(new UnrecoverableDurableExecutionException( + ErrorObject.builder().errorMessage("probe cleanup").build(), true)); + } catch (UnrecoverableDurableExecutionException expected) { + } + } + caller.shutdownNow(); + workers.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(5, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new AssertionError(interrupted); + } + } + + private static DurableExecutionInput input(List operations) { + var all = new ArrayList(); + all.add(Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build()); + all.addAll(operations); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/step-retry/execution", + "token", + CheckpointUpdatedExecutionState.builder().operations(all).build()); + } + + @ParameterizedTest + @CsvSource({"true,false", "false,true", "true,true"}) + void liveRetryRetiresItsWorkerBeforeReplacementWithoutLosingActivity( + boolean holdWorkerExit, boolean holdExecutorReturn) throws Exception { + var gate = new HandoffGate(holdWorkerExit, holdExecutorReturn); + var client = new DelayedReadyClient(gate); + var calls = new AtomicInteger(); + var caller = Executors.newSingleThreadExecutor(); + var config = DurableConfig.builder() + .withExecutorService(gate) + .withDurableExecutionClient(client) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO) + .build(); + var retry = StepConfig.builder() + .retryStrategy((error, attempt) -> + attempt == 1 ? RetryDecision.retry(Duration.ofSeconds(1)) : RetryDecision.fail()) + .build(); + BiFunction handler = (value, context) -> { + var step = context.stepAsync( + "retry-step", + String.class, + stepContext -> { + if (calls.incrementAndGet() == 1) { + gate.firstWorker.set(Thread.currentThread()); + throw new IllegalStateException("retry once"); + } + gate.nextCheckEntered.countDown(); + await(gate.releaseNextCheck); + return "stored"; + }, + retry); + // Keep the root active until the live retry has registered its backend poll. + await(gate.pollEntered); + return step.get(); + }; + try (var mdc = new MdcClearGate(gate)) { + try { + var result = caller.submit(() -> + DurableExecutor.execute(input(List.of()), null, TypeToken.get(String.class), handler, config)); + await(gate.pollEntered); + if (holdWorkerExit) await(gate.workerExitEntered); + else await(gate.executorReturnEntered); + assertTrue(client.advanceTime()); + gate.releaseReadyResponse.countDown(); + await(gate.readyResponseReturned); + awaitCheckpointHandoff(gate.pollThread.get()); + assertThrows(TimeoutException.class, () -> result.get(100, TimeUnit.MILLISECONDS)); + assertEquals(1, calls.get()); + gate.releaseWorkerExit.countDown(); + if (holdExecutorReturn) await(gate.executorReturnEntered); + gate.releaseExecutorReturn.countDown(); + await(gate.nextCheckEntered); + assertThrows( + TimeoutException.class, + () -> result.get(100, TimeUnit.MILLISECONDS), + "The replacement retains activity after its previous worker deregisters"); + gate.releaseNextCheck.countDown(); + var completed = result.get(5, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.SUCCEEDED, completed.status()); + assertEquals("\"stored\"", completed.result()); + assertEquals( + OperationStatus.SUCCEEDED, + client.getOperationByName("retry-step").status()); + var replay = caller.submit(() -> DurableExecutor.execute( + input(client.getAllOperations()), null, TypeToken.get(String.class), handler, config)); + assertEquals(completed, replay.get(5, TimeUnit.SECONDS)); + assertEquals(2, calls.get(), "The completed retry body is skipped on replay"); + } finally { + gate.releaseAll(); + caller.shutdownNow(); + gate.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(gate.awaitTermination(5, TimeUnit.SECONDS)); + } + } + } + + private static void awaitCheckpointHandoff(Thread checkpoint) { + var deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(3); + while (System.nanoTime() < deadline) { + for (var entry : Thread.getAllStackTraces().entrySet()) { + var continuation = entry.getKey(); + var stack = entry.getValue(); + if (continuation.getState() != Thread.State.WAITING + || !continuation.getName().startsWith("durable-sdk-internal-") + || Arrays.stream(stack) + .noneMatch(frame -> frame.getClassName().endsWith("StepOperation")) + || Arrays.stream(stack) + .noneMatch(frame -> frame.getClassName().equals(CompletableFuture.class.getName()) + && frame.getMethodName().equals("join"))) continue; + assertNotSame(checkpoint, continuation, "Worker publication must not block the checkpoint callback"); + assertTrue(Arrays.stream(stack) + .noneMatch(frame -> frame.getClassName().endsWith("ApiRequestDelayedBatcher"))); + if (Arrays.stream(checkpoint.getStackTrace()) + .anyMatch(frame -> frame.getClassName().endsWith("CheckpointManager") + && frame.getMethodName().equals("checkpointBatch"))) continue; + var bean = ManagementFactory.getThreadMXBean(); + if (bean.isObjectMonitorUsageSupported()) { + var info = bean.getThreadInfo(new long[] {continuation.getId()}, true, true)[0]; + assertEquals(0, info.getLockedMonitors().length); + System.out.println("READY_HANDOFF continuation=" + continuation.getName() + " checkpoint=" + + checkpoint.getName() + " continuationMonitors=[] checkpointBatchReturned=true"); + } + return; + } + LockSupport.parkNanos(TimeUnit.MILLISECONDS.toNanos(1)); + } + fail("The independent READY continuation did not reach the controlled old-worker handoff window"); + } + + private static final class DelayedReadyClient extends LocalMemoryExecutionClient { + private final HandoffGate gate; + private final AtomicBoolean firstPoll = new AtomicBoolean(true); + + private DelayedReadyClient(HandoffGate gate) { + this.gate = gate; + } + + @Override + public CheckpointDurableExecutionResponse checkpoint(String arn, String token, List updates) { + if (updates.isEmpty() && firstPoll.compareAndSet(true, false)) { + gate.pollThread.set(Thread.currentThread()); + gate.pollEntered.countDown(); + await(gate.releaseReadyResponse); + var response = super.checkpoint(arn, token, updates); + gate.readyResponseReturned.countDown(); + return response; + } + return super.checkpoint(arn, token, updates); + } + } + + private static final class HandoffGate extends AbstractExecutorService { + private final ExecutorService delegate = Executors.newCachedThreadPool(); + private final AtomicInteger submitted = new AtomicInteger(); + private final AtomicBoolean exitHeld = new AtomicBoolean(); + private final AtomicReference firstWorker = new AtomicReference<>(); + private final AtomicReference pollThread = new AtomicReference<>(); + private final CountDownLatch pollEntered = new CountDownLatch(1); + private final CountDownLatch releaseReadyResponse = new CountDownLatch(1); + private final CountDownLatch readyResponseReturned = new CountDownLatch(1); + private final CountDownLatch workerExitEntered = new CountDownLatch(1); + private final CountDownLatch releaseWorkerExit = new CountDownLatch(1); + private final CountDownLatch executorReturnEntered = new CountDownLatch(1); + private final CountDownLatch releaseExecutorReturn = new CountDownLatch(1); + private final CountDownLatch nextCheckEntered = new CountDownLatch(1); + private final CountDownLatch releaseNextCheck = new CountDownLatch(1); + private final boolean holdWorkerExit; + private final boolean holdExecutorReturn; + + private HandoffGate(boolean holdWorkerExit, boolean holdExecutorReturn) { + this.holdWorkerExit = holdWorkerExit; + this.holdExecutorReturn = holdExecutorReturn; + } + + @Override + public void execute(Runnable task) { + var number = submitted.incrementAndGet(); + var finished = new CountDownLatch(1); + delegate.execute(() -> { + try { + task.run(); + } finally { + finished.countDown(); + } + }); + // Completing work before execute() returns is permitted by Executor's contract. + if (holdExecutorReturn && number == 2) { + await(finished); + executorReturnEntered.countDown(); + await(releaseExecutorReturn); + } + } + + private void releaseAll() { + releaseReadyResponse.countDown(); + releaseWorkerExit.countDown(); + releaseExecutorReturn.countDown(); + releaseNextCheck.countDown(); + } + + @Override + public void shutdown() { + delegate.shutdown(); + } + + @Override + public List shutdownNow() { + releaseAll(); + return delegate.shutdownNow(); + } + + @Override + public boolean isShutdown() { + return delegate.isShutdown(); + } + + @Override + public boolean isTerminated() { + return delegate.isTerminated(); + } + + @Override + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(timeout, unit); + } + } + + private static final class MdcClearGate implements AutoCloseable { + private final MDCAdapter previous = MDC.getMDCAdapter(); + private final Method setter = MDC.class.getDeclaredMethod("setMDCAdapter", MDCAdapter.class); + + private MdcClearGate(HandoffGate gate) throws Exception { + setter.setAccessible(true); + setter.invoke(null, new MDCAdapter() { + public void clear() { + if (gate.holdWorkerExit + && Thread.currentThread() == gate.firstWorker.get() + && gate.exitHeld.compareAndSet(false, true)) { + gate.workerExitEntered.countDown(); + await(gate.releaseWorkerExit); + } + previous.clear(); + } + + public void put(String key, String value) { + previous.put(key, value); + } + + public String get(String key) { + return previous.get(key); + } + + public void remove(String key) { + previous.remove(key); + } + + public Map getCopyOfContextMap() { + return previous.getCopyOfContextMap(); + } + + public void setContextMap(Map values) { + previous.setContextMap(values); + } + + public void pushByKey(String key, String value) { + previous.pushByKey(key, value); + } + + public String popByKey(String key) { + return previous.popByKey(key); + } + + public Deque getCopyOfDequeByKey(String key) { + return previous.getCopyOfDequeByKey(key); + } + + public void clearDequeByKey(String key) { + previous.clearDequeByKey(key); + } + }); + } + + public void close() throws Exception { + setter.invoke(null, previous); + } + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/WaitForConditionReadinessIntegrationTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/WaitForConditionReadinessIntegrationTest.java new file mode 100644 index 000000000..9ed8cd295 --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/WaitForConditionReadinessIntegrationTest.java @@ -0,0 +1,672 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.lang.management.ManagementFactory; +import java.lang.reflect.Method; +import java.nio.file.Files; +import java.nio.file.Path; +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Deque; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.*; +import java.util.concurrent.atomic.*; +import java.util.concurrent.locks.LockSupport; +import java.util.function.BiFunction; +import java.util.function.IntFunction; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.ValueSource; +import org.slf4j.MDC; +import org.slf4j.spi.MDCAdapter; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.config.WaitForConditionConfig; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.DurableExecutionOutput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.model.WaitForConditionResult; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class WaitForConditionReadinessIntegrationTest { + @TempDir + Path subprocessLogs; + + @ParameterizedTest + @ValueSource(strings = {"async", "direct", "submit-and-wait"}) + void delayedReadyResumesWithSynchronousOperationExecutors(String mode) throws Exception { + assertProbe(SynchronousProbe.class, mode, "SYNC_READY_SUCCESS checks=2 replayChecks=2"); + } + + @ParameterizedTest + @ValueSource(strings = {"cached", "bounded"}) + void immediateReadyYieldsToTheSiblingThatUnlocksItsCondition(String mode) throws Exception { + assertProbe(FairnessProbe.class, mode, "FAIR_READY_SUCCESS siblingCalls=1 replayStable=true"); + } + + private void assertProbe(Class probe, String mode, String successMarker) throws Exception { + var log = subprocessLogs.resolve(probe.getSimpleName() + "-" + mode + ".log"); + var process = new ProcessBuilder( + Path.of(System.getProperty("java.home"), "bin", "java").toString(), + "-cp", + System.getProperty("java.class.path"), + probe.getName(), + mode) + .redirectErrorStream(true) + .redirectOutput(log.toFile()) + .start(); + try { + assertTrue(process.waitFor(15, TimeUnit.SECONDS), "The isolated READY probe exceeded its cleanup bound"); + var output = Files.readString(log); + System.out.println("SYNC_READY_PROBE mode=" + mode + " exit=" + process.exitValue() + "\n" + output); + assertEquals(0, process.exitValue(), output); + assertTrue(output.contains(successMarker), output); + } finally { + if (process.isAlive()) process.destroyForcibly(); + assertTrue(process.waitFor(5, TimeUnit.SECONDS)); + } + } + + public static final class FairnessProbe { + public static void main(String[] args) throws Exception { + var workers = args[0].equals("bounded") ? Executors.newFixedThreadPool(2) : Executors.newCachedThreadPool(); + var siblingQueued = new CountDownLatch(1); + var unlocked = new AtomicBoolean(); + var checks = new AtomicInteger(); + var siblings = new AtomicInteger(); + var run = new Run(new ImmediateReadyClient(), workers, value -> { + throw new AssertionError("The explicit public handler supplies its own condition"); + }); + run.customHandler = (input, context) -> { + var condition = context.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + checks.incrementAndGet(); + await(siblingQueued); + return unlocked.get() + ? WaitForConditionResult.stopPolling(state) + : WaitForConditionResult.continuePolling(state + 1); + }, + run.wait); + var sibling = context.stepAsync("unlock-condition", String.class, step -> { + siblings.incrementAndGet(); + unlocked.set(true); + return "unlocked"; + }); + siblingQueued.countDown(); + var result = condition.get(); + assertEquals("unlocked", sibling.get()); + return String.valueOf(result); + }; + try { + var first = run.start().get(3, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.SUCCEEDED, first.status()); + assertEquals(1, siblings.get()); + var completedChecks = checks.get(); + var checkpoints = run.client.getOperationUpdates().size(); + assertEquals(first, run.start().get(3, TimeUnit.SECONDS)); + assertEquals(completedChecks, checks.get()); + assertEquals(1, siblings.get()); + assertEquals(checkpoints, run.client.getOperationUpdates().size()); + System.out.println("FAIR_READY_SUCCESS siblingCalls=1 replayStable=true checks=" + completedChecks); + run.close(); + } catch (Throwable failure) { + System.out.println("FAIR_READY_FAILURE mode=" + args[0] + " checks=" + checks.get() + " siblingCalls=" + + siblings.get() + " queued=" + + ((ThreadPoolExecutor) workers).getQueue().size()); + failure.printStackTrace(System.out); + Thread.getAllStackTraces().forEach((thread, stack) -> { + if (!thread.getName().startsWith("durable-sdk-internal") + && !thread.getName().startsWith("pool-")) return; + System.out.println("THREAD " + thread.getName() + " " + thread.getState()); + Arrays.stream(stack).forEach(frame -> System.out.println(" " + frame)); + }); + System.exit(2); + } + } + } + + /** A separate process bounds cleanup of an old-code deadlock without altering the observed protocol. */ + public static final class SynchronousProbe { + public static void main(String[] args) throws Exception { + var gate = new HandoffGate(true, false); + var workers = new SynchronousContinuationExecutor(args[0]); + var client = new DelayedReadyClient(gate); + var run = new Run(client, workers, value -> { + if (value == 1) gate.firstWorker.set(Thread.currentThread()); + return value == 2 + ? WaitForConditionResult.stopPolling(value) + : WaitForConditionResult.continuePolling(2); + }); + try (var mdc = new MdcClearGate(gate)) { + var first = run.start(); + await(gate.pollEntered); + await(gate.workerExitEntered); + assertTrue(client.advanceTime()); + gate.releaseReadyResponse.countDown(); + await(gate.readyResponseReturned); + gate.releaseWorkerExit.countDown(); + assertEquals( + ExecutionStatus.SUCCEEDED, + first.get(3, TimeUnit.SECONDS).status()); + assertEquals(2, run.checks.get()); + assertEquals( + ExecutionStatus.SUCCEEDED, + run.start().get(3, TimeUnit.SECONDS).status()); + assertEquals(2, run.checks.get()); + assertTrue(workers.continuationDispatches.get() > 0); + System.out.println("SYNC_READY_SUCCESS checks=2 replayChecks=2 mode=" + args[0]); + run.close(); + gate.shutdownNow(); + } catch (Throwable failure) { + System.out.println("SYNC_READY_BLOCKED mode=" + args[0] + " failure=" + + failure.getClass().getName()); + failure.printStackTrace(System.out); + var bean = ManagementFactory.getThreadMXBean(); + Thread.getAllStackTraces().forEach((thread, stack) -> { + if (!thread.getName().startsWith("durable-sdk-internal") + && !thread.getName().startsWith("sync-ready-worker")) return; + var info = bean.getThreadInfo(new long[] {thread.getId()}, true, true)[0]; + System.out.println("THREAD " + thread.getName() + " " + thread.getState() + " monitors=" + + (info == null ? "[]" : Arrays.toString(info.getLockedMonitors()))); + Arrays.stream(stack).forEach(frame -> System.out.println(" " + frame)); + }); + // The probe owns no cloud resources; exiting the child JVM releases only this fixture's blocked + // threads. + System.exit(2); + } + } + } + + private static final class SynchronousContinuationExecutor extends AbstractExecutorService { + private final String mode; + private final AtomicInteger submissions = new AtomicInteger(); + private final AtomicInteger continuationDispatches = new AtomicInteger(); + private final ExecutorService delegate = Executors.newCachedThreadPool(task -> { + var thread = new Thread(task, "sync-ready-worker"); + thread.setDaemon(true); + return thread; + }); + + private SynchronousContinuationExecutor(String mode) { + this.mode = mode; + } + + @Override + public void execute(Runnable task) { + if (submissions.incrementAndGet() <= 2) { + delegate.execute(task); // Root handler and first check retain their normal asynchronous setup. + return; + } + continuationDispatches.incrementAndGet(); + if (mode.equals("direct")) task.run(); + else if (mode.equals("submit-and-wait")) { + try { + delegate.submit(task).get(); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } catch (ExecutionException failure) { + throw new AssertionError(failure.getCause()); + } + } else delegate.execute(task); + } + + @Override + public void shutdown() { + delegate.shutdown(); + } + + @Override + public List shutdownNow() { + return delegate.shutdownNow(); + } + + @Override + public boolean isShutdown() { + return delegate.isShutdown(); + } + + @Override + public boolean isTerminated() { + return delegate.isTerminated(); + } + + @Override + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(timeout, unit); + } + } + + @ParameterizedTest + @ValueSource(ints = {2, 25}) + void readyInRetryResponseContinuesWithoutSuspensionAndReplaysStoredResult(int threshold) throws Exception { + try (var run = new Run( + new ImmediateReadyClient(), + Executors.newCachedThreadPool(), + value -> value >= threshold + ? WaitForConditionResult.stopPolling(value) + : WaitForConditionResult.continuePolling(value + 1))) { + var first = run.start().get(5, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.SUCCEEDED, first.status(), "READY work must not produce an invalid PENDING"); + assertEquals("\"" + threshold + "\"", first.result()); + assertEquals(threshold, run.checks.get()); + var updates = run.client.getOperationUpdates().size(); + var replay = run.start().get(5, TimeUnit.SECONDS); + assertEquals(first, replay); + assertEquals(threshold, run.checks.get(), "Completed checks must not be repeated on replay"); + assertEquals(updates, run.client.getOperationUpdates().size()); + } + } + + @Test + void immediateReadyFailureRemainsCheckpointedOnReplay() throws Exception { + try (var run = new Run(new ImmediateReadyClient(), Executors.newCachedThreadPool(), value -> { + if (value == 2) throw new IllegalStateException("predicate failure"); + return WaitForConditionResult.continuePolling(value + 1); + })) { + var first = run.start().get(5, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.FAILED, first.status()); + assertEquals("predicate failure", first.error().errorMessage()); + var updates = run.client.getOperationUpdates().size(); + var replay = run.start().get(5, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.FAILED, replay.status()); + assertEquals(first.error().errorType(), replay.error().errorType()); + assertEquals(first.error().errorMessage(), replay.error().errorMessage()); + assertEquals(2, run.checks.get()); + assertEquals(updates, run.client.getOperationUpdates().size()); + } + } + + @Test + void pendingRetryStillSuspendsAndResumesWhenBackendBecomesReady() throws Exception { + try (var run = new Run( + new LocalMemoryExecutionClient(), + Executors.newCachedThreadPool(), + value -> value == 2 + ? WaitForConditionResult.stopPolling(value) + : WaitForConditionResult.continuePolling(value + 1))) { + var first = run.start().get(5, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.PENDING, first.status()); + assertEquals( + OperationStatus.PENDING, + run.client.getOperationByName("condition").status()); + assertEquals(1, run.checks.get()); + assertTrue(run.client.advanceTime()); + var resumed = run.start().get(5, TimeUnit.SECONDS); + assertEquals(ExecutionStatus.SUCCEEDED, resumed.status()); + assertEquals("\"2\"", resumed.result()); + assertEquals(2, run.checks.get()); + assertEquals(resumed, run.start().get(5, TimeUnit.SECONDS)); + assertEquals(2, run.checks.get()); + } + } + + @ParameterizedTest + @CsvSource({"true,false", "false,true", "true,true"}) + void asynchronousReadyHandoffDoesNotDeadlockOrLoseItsActivityLease( + boolean holdWorkerExit, boolean holdExecutorReturn) throws Exception { + var gate = new HandoffGate(holdWorkerExit, holdExecutorReturn); + var client = new DelayedReadyClient(gate); + try (var mdc = new MdcClearGate(gate); + var run = new Run(client, gate, value -> { + if (value == 1) { + gate.firstWorker.set(Thread.currentThread()); + return WaitForConditionResult.continuePolling(2); + } + gate.nextCheckEntered.countDown(); + await(gate.releaseNextCheck); + return WaitForConditionResult.stopPolling(value); + })) { + try { + var result = run.start(); + await(gate.pollEntered); + if (holdWorkerExit) await(gate.workerExitEntered); + else await(gate.executorReturnEntered); + assertTrue(client.advanceTime()); + gate.releaseReadyResponse.countDown(); + await(gate.readyResponseReturned); + awaitCheckpointHandoff(gate.pollThread.get()); + assertThrows(TimeoutException.class, () -> result.get(100, TimeUnit.MILLISECONDS)); + assertEquals(1, run.checks.get()); + gate.releaseWorkerExit.countDown(); + if (holdExecutorReturn) await(gate.executorReturnEntered); + gate.releaseExecutorReturn.countDown(); + await(gate.nextCheckEntered); + assertThrows( + TimeoutException.class, + () -> result.get(100, TimeUnit.MILLISECONDS), + "The next check must stay active after the old worker deregisters and the checkpoint returns"); + gate.releaseNextCheck.countDown(); + assertEquals( + ExecutionStatus.SUCCEEDED, + result.get(5, TimeUnit.SECONDS).status()); + assertEquals(2, run.checks.get()); + assertEquals(4, run.stateReads.get(), "In-process continuation must reuse normalized state"); + assertEquals( + ExecutionStatus.SUCCEEDED, + run.start().get(5, TimeUnit.SECONDS).status()); + assertEquals(2, run.checks.get()); + assertEquals(5, run.stateReads.get(), "Completed replay only decodes the stored result"); + } finally { + gate.releaseAll(); + } + } + } + + private static void awaitCheckpointHandoff(Thread checkpoint) { + var deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(3); + while (System.nanoTime() < deadline) { + for (var entry : Thread.getAllStackTraces().entrySet()) { + var continuation = entry.getKey(); + var stack = entry.getValue(); + if (continuation.getState() != Thread.State.WAITING + || !continuation.getName().startsWith("durable-sdk-internal-") + || Arrays.stream(stack) + .noneMatch(frame -> frame.getClassName().endsWith("WaitForConditionOperation")) + || Arrays.stream(stack) + .noneMatch(frame -> frame.getClassName().equals(CompletableFuture.class.getName()) + && frame.getMethodName().equals("join"))) continue; + assertNotSame(checkpoint, continuation, "Worker publication must not block the checkpoint callback"); + assertTrue(Arrays.stream(stack) + .noneMatch(frame -> frame.getClassName().endsWith("ApiRequestDelayedBatcher"))); + if (Arrays.stream(checkpoint.getStackTrace()) + .anyMatch(frame -> frame.getClassName().endsWith("CheckpointManager") + && frame.getMethodName().equals("checkpointBatch"))) continue; + var bean = ManagementFactory.getThreadMXBean(); + if (bean.isObjectMonitorUsageSupported()) { + var info = bean.getThreadInfo(new long[] {continuation.getId()}, true, true)[0]; + assertEquals(0, info.getLockedMonitors().length); + System.out.println("READY_HANDOFF continuation=" + continuation.getName() + " checkpoint=" + + checkpoint.getName() + " continuationMonitors=[] checkpointBatchReturned=true"); + } + return; + } + LockSupport.parkNanos(TimeUnit.MILLISECONDS.toNanos(1)); + } + fail("The independent READY continuation did not reach the controlled old-worker handoff window"); + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(5, TimeUnit.SECONDS), "Controlled readiness gate was not released"); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throw new AssertionError(error); + } + } + + private static final class ImmediateReadyClient extends LocalMemoryExecutionClient { + @Override + public CheckpointDurableExecutionResponse checkpoint(String arn, String token, List updates) { + var applied = super.checkpoint(arn, token, updates); + // Advance the test backend's virtual time before materializing the response. The configured + // one-second delay is valid under Lambda's minimum; no zero-delay service request is assumed. + advanceTime(); + var ready = super.checkpoint(arn, applied.checkpointToken(), List.of()); + var states = new LinkedHashMap(); + applied.newExecutionState().operations().forEach(op -> states.put(op.id(), op)); + ready.newExecutionState().operations().forEach(op -> states.put(op.id(), op)); + return ready.toBuilder() + .newExecutionState(CheckpointUpdatedExecutionState.builder() + .operations(states.values()) + .build()) + .build(); + } + } + + private static final class DelayedReadyClient extends LocalMemoryExecutionClient { + private final HandoffGate gate; + private final AtomicBoolean firstPoll = new AtomicBoolean(true); + + private DelayedReadyClient(HandoffGate gate) { + this.gate = gate; + } + + @Override + public CheckpointDurableExecutionResponse checkpoint(String arn, String token, List updates) { + if (updates.isEmpty() && firstPoll.compareAndSet(true, false)) { + gate.pollThread.set(Thread.currentThread()); + gate.pollEntered.countDown(); + await(gate.releaseReadyResponse); + var response = super.checkpoint(arn, token, updates); + gate.readyResponseReturned.countDown(); + return response; + } + return super.checkpoint(arn, token, updates); + } + } + + private static final class HandoffGate extends AbstractExecutorService { + private final ExecutorService delegate = Executors.newCachedThreadPool(); + private final AtomicInteger submitted = new AtomicInteger(); + private final AtomicBoolean exitHeld = new AtomicBoolean(); + private final AtomicReference firstWorker = new AtomicReference<>(); + private final AtomicReference pollThread = new AtomicReference<>(); + private final CountDownLatch pollEntered = new CountDownLatch(1); + private final CountDownLatch releaseReadyResponse = new CountDownLatch(1); + private final CountDownLatch readyResponseReturned = new CountDownLatch(1); + private final CountDownLatch workerExitEntered = new CountDownLatch(1); + private final CountDownLatch releaseWorkerExit = new CountDownLatch(1); + private final CountDownLatch executorReturnEntered = new CountDownLatch(1); + private final CountDownLatch releaseExecutorReturn = new CountDownLatch(1); + private final CountDownLatch nextCheckEntered = new CountDownLatch(1); + private final CountDownLatch releaseNextCheck = new CountDownLatch(1); + private final boolean holdWorkerExit; + private final boolean holdExecutorReturn; + + private HandoffGate(boolean holdWorkerExit, boolean holdExecutorReturn) { + this.holdWorkerExit = holdWorkerExit; + this.holdExecutorReturn = holdExecutorReturn; + } + + @Override + public void execute(Runnable task) { + var number = submitted.incrementAndGet(); + var finished = new CountDownLatch(1); + delegate.execute(() -> { + try { + task.run(); + } finally { + finished.countDown(); + } + }); + // Completing work before execute() returns is permitted by Executor's contract. + if (holdExecutorReturn && number == 2) { + await(finished); + executorReturnEntered.countDown(); + await(releaseExecutorReturn); + } + } + + private void releaseAll() { + releaseReadyResponse.countDown(); + releaseWorkerExit.countDown(); + releaseExecutorReturn.countDown(); + releaseNextCheck.countDown(); + } + + @Override + public void shutdown() { + delegate.shutdown(); + } + + @Override + public List shutdownNow() { + releaseAll(); + return delegate.shutdownNow(); + } + + @Override + public boolean isShutdown() { + return delegate.isShutdown(); + } + + @Override + public boolean isTerminated() { + return delegate.isTerminated(); + } + + @Override + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return delegate.awaitTermination(timeout, unit); + } + } + + private static final class MdcClearGate implements AutoCloseable { + private final MDCAdapter previous = MDC.getMDCAdapter(); + private final Method setter = MDC.class.getDeclaredMethod("setMDCAdapter", MDCAdapter.class); + + private MdcClearGate(HandoffGate gate) throws Exception { + setter.setAccessible(true); + setter.invoke(null, new MDCAdapter() { + public void clear() { + if (gate.holdWorkerExit + && Thread.currentThread() == gate.firstWorker.get() + && gate.exitHeld.compareAndSet(false, true)) { + gate.workerExitEntered.countDown(); + await(gate.releaseWorkerExit); + } + previous.clear(); + } + + public void put(String key, String value) { + previous.put(key, value); + } + + public String get(String key) { + return previous.get(key); + } + + public void remove(String key) { + previous.remove(key); + } + + public Map getCopyOfContextMap() { + return previous.getCopyOfContextMap(); + } + + public void setContextMap(Map values) { + previous.setContextMap(values); + } + + public void pushByKey(String key, String value) { + previous.pushByKey(key, value); + } + + public String popByKey(String key) { + return previous.popByKey(key); + } + + public Deque getCopyOfDequeByKey(String key) { + return previous.getCopyOfDequeByKey(key); + } + + public void clearDequeByKey(String key) { + previous.clearDequeByKey(key); + } + }); + } + + public void close() throws Exception { + setter.invoke(null, previous); + } + } + + private static final class Run implements AutoCloseable { + private final LocalMemoryExecutionClient client; + private final ExecutorService workers; + private final ExecutorService caller = Executors.newSingleThreadExecutor(); + private final AtomicInteger checks = new AtomicInteger(); + private final AtomicInteger stateReads = new AtomicInteger(); + private final SerDes stateSerDes = new SerDes() { + private final JacksonSerDes delegate = new JacksonSerDes(); + + public String serialize(Object value) { + return delegate.serialize(value); + } + + public T deserialize(String value, TypeToken type) { + stateReads.incrementAndGet(); + return delegate.deserialize(value, type); + } + }; + private final IntFunction> check; + private BiFunction customHandler; + private final DurableConfig config; + private final WaitForConditionConfig wait = WaitForConditionConfig.builder() + .initialState(1) + .serDes(stateSerDes) + .waitStrategy((value, attempt) -> Duration.ofSeconds(1)) + .build(); + + private Run( + LocalMemoryExecutionClient client, + ExecutorService workers, + IntFunction> check) { + this.client = client; + this.workers = workers; + this.check = check; + config = DurableConfig.builder() + .withExecutorService(workers) + .withDurableExecutionClient(client) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO) + .build(); + } + + private Future start() { + var execution = Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + var history = new ArrayList<>(client.getAllOperations()); + history.add(0, execution); + var input = new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/ready/execution", + "token", + CheckpointUpdatedExecutionState.builder() + .operations(history) + .build()); + return caller.submit(() -> DurableExecutor.execute( + input, + null, + TypeToken.get(String.class), + customHandler != null + ? customHandler + : (value, ctx) -> String.valueOf(ctx.waitForCondition( + "condition", + Integer.class, + (state, step) -> { + checks.incrementAndGet(); + return check.apply(state); + }, + wait)), + config)); + } + + @Override + public void close() throws Exception { + caller.shutdownNow(); + workers.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/WaitForConditionTerminalReplayIntegrationTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/WaitForConditionTerminalReplayIntegrationTest.java new file mode 100644 index 000000000..ed11e1d60 --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/WaitForConditionTerminalReplayIntegrationTest.java @@ -0,0 +1,221 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import software.amazon.awssdk.services.lambda.model.CheckpointDurableExecutionResponse; +import software.amazon.awssdk.services.lambda.model.CheckpointUpdatedExecutionState; +import software.amazon.awssdk.services.lambda.model.ErrorObject; +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.awssdk.services.lambda.model.StepDetails; +import software.amazon.lambda.durable.config.WaitForConditionConfig; +import software.amazon.lambda.durable.exception.WaitForConditionFailedException; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.DurableExecutionOutput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.model.WaitForConditionResult; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class WaitForConditionTerminalReplayIntegrationTest { + private static final String CACHED_ERROR = "cached terminal error"; + + static Stream terminals() { + return Stream.of( + OperationStatus.FAILED, + OperationStatus.CANCELLED, + OperationStatus.TIMED_OUT, + OperationStatus.STOPPED) + .flatMap(status -> Stream.of("absent", "empty", "stored") + .flatMap(details -> Stream.of("replay", "retry-response", "poll-response") + .map(delivery -> Arguments.of(status, details, delivery)))); + } + + @ParameterizedTest + @MethodSource("terminals") + void terminalSnapshotsReplayWithoutRepeatingChecksOrCheckpoints( + OperationStatus status, String details, String delivery) throws Exception { + var deliverWhileRunning = !delivery.equals("replay"); + var retried = new AtomicBoolean(); + var delivered = new AtomicBoolean(); + var pollEntered = new CountDownLatch(1); + var releasePoll = new CountDownLatch(1); + var terminal = new AtomicReference(); + var client = new LocalMemoryExecutionClient() { + @Override + public CheckpointDurableExecutionResponse checkpoint( + String arn, String token, List updates) { + var response = super.checkpoint(arn, token, updates); + var retry = updates.stream().anyMatch(u -> u.action() == OperationAction.RETRY); + if (retry) retried.set(true); + var returnTerminal = delivery.equals("retry-response") && retry + || delivery.equals("poll-response") && retried.get() && updates.isEmpty(); + if (!returnTerminal || !delivered.compareAndSet(false, true)) return response; + if (delivery.equals("poll-response")) { + pollEntered.countDown(); + await(releasePoll); + } + // Exercise an actual checkpoint-completion boundary with the backend's terminal snapshot. + var condition = getOperationByName("condition"); + terminal.set(terminalSnapshot(condition, status, details)); + return response.toBuilder() + .newExecutionState(CheckpointUpdatedExecutionState.builder() + .operations(terminal.get()) + .build()) + .build(); + } + }; + var workers = Executors.newCachedThreadPool(); + var caller = Executors.newSingleThreadExecutor(); + var checks = new AtomicInteger(); + try { + var config = DurableConfig.builder() + .withExecutorService(workers) + .withDurableExecutionClient(client) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO) + .build(); + var first = caller.submit(() -> execute(List.of(), config, checks, status, details, () -> { + if (delivery.equals("poll-response")) { + await(pollEntered); + releasePoll.countDown(); + } + })) + .get(5, TimeUnit.SECONDS); + if (deliverWhileRunning) assertHandled(first, status, details); + else { + assertEquals(ExecutionStatus.PENDING, first.status()); + terminal.set(terminalSnapshot(client.getOperationByName("condition"), status, details)); + } + assertEquals(1, checks.get()); + var updates = client.getOperationUpdates().size(); + var history = List.of(terminal.get()); + var replay = caller.submit(() -> execute(history, config, checks, status, details, () -> {})) + .get(5, TimeUnit.SECONDS); + assertHandled(replay, status, details); + var repeated = caller.submit(() -> execute(history, config, checks, status, details, () -> {})) + .get(5, TimeUnit.SECONDS); + assertEquals(replay, repeated); + assertEquals(1, checks.get(), "Neither cancelled polling nor completed replay may re-enter the predicate"); + assertEquals(updates, client.getOperationUpdates().size(), "Replay must not rewrite a terminal checkpoint"); + } finally { + releasePoll.countDown(); + caller.shutdownNow(); + workers.shutdownNow(); + assertTrue(caller.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(workers.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(5, TimeUnit.SECONDS), "Controlled terminal poll was not released"); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + + private static DurableExecutionOutput execute( + List history, + DurableConfig config, + AtomicInteger checks, + OperationStatus status, + String details, + Runnable beforeAwait) { + var operations = new ArrayList(); + operations.add(Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build()); + operations.addAll(history); + var input = new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/terminal/execution", + "token", + CheckpointUpdatedExecutionState.builder().operations(operations).build()); + return DurableExecutor.execute( + input, + null, + TypeToken.get(String.class), + (value, context) -> { + try { + var future = context.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + checks.incrementAndGet(); + return WaitForConditionResult.continuePolling(1); + }, + WaitForConditionConfig.builder() + .initialState(0) + .waitStrategy((state, attempt) -> Duration.ofSeconds(1)) + .build()); + beforeAwait.run(); + future.get(); + return "unexpected success"; + } catch (WaitForConditionFailedException failure) { + assertNotEquals("stored", details); + assertEquals(status, failure.getOperationStatus()); + assertEquals("condition", failure.getOperation().name()); + return status + ":no error details"; + } catch (IllegalStateException failure) { + assertEquals("stored", details); + assertEquals(CACHED_ERROR, failure.getMessage()); + return CACHED_ERROR; + } + }, + config); + } + + private static void assertHandled(DurableExecutionOutput output, OperationStatus status, String details) { + assertEquals( + ExecutionStatus.SUCCEEDED, + output.status(), + () -> output.error() == null + ? "No error payload" + : output.error().errorType() + ": " + output.error().errorMessage()); + assertEquals( + new JacksonSerDes().serialize(details.equals("stored") ? CACHED_ERROR : status + ":no error details"), + output.result()); + } + + private static Operation terminalSnapshot(Operation pending, OperationStatus status, String details) { + StepDetails snapshot = null; + if (details.equals("empty")) snapshot = StepDetails.builder().attempt(1).build(); + else if (details.equals("stored")) + snapshot = StepDetails.builder() + .attempt(1) + .error(ErrorObject.builder() + .errorType(IllegalStateException.class.getName()) + .errorMessage(CACHED_ERROR) + .errorData(new JacksonSerDes().serialize(new IllegalStateException(CACHED_ERROR))) + .build()) + .build(); + return pending.toBuilder().status(status).stepDetails(snapshot).build(); + } +} diff --git a/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/execution/ContinuationShutdownIntegrationTest.java b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/execution/ContinuationShutdownIntegrationTest.java new file mode 100644 index 000000000..94b3489f8 --- /dev/null +++ b/sdk-integration-tests/src/test/java/software/amazon/lambda/durable/execution/ContinuationShutdownIntegrationTest.java @@ -0,0 +1,431 @@ +// 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 java.time.Duration; +import java.time.Instant; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.concurrent.*; +import java.util.concurrent.atomic.*; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.config.CompletionConfig; +import software.amazon.lambda.durable.config.ParallelConfig; +import software.amazon.lambda.durable.config.WaitForConditionConfig; +import software.amazon.lambda.durable.context.BaseContext; +import software.amazon.lambda.durable.context.BaseContextImpl; +import software.amazon.lambda.durable.context.DurableContextImpl; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.WaitForConditionResult; +import software.amazon.lambda.durable.operation.BaseDurableOperation; +import software.amazon.lambda.durable.testing.local.LocalMemoryExecutionClient; + +class ContinuationShutdownIntegrationTest { + private ControlledManager initialManager; + private BaseContext previousContext; + private ThreadContext previousThreadContext; + + private void rememberContext(ControlledManager manager) { + if (initialManager != null) return; + initialManager = manager; + previousContext = BaseContext.getCurrentContext(); + previousThreadContext = manager.getCurrentThreadContext(); + } + + @AfterEach + void restoreCallerContext() { + if (initialManager == null) return; + BaseContextImpl.setCurrentContext(previousContext); + initialManager.setCurrentThreadContext(previousThreadContext); + } + + @Test + void runningPredicateFinishesItsCheckpointAndReplaysWithoutRepeatingWork() throws Exception { + var closed = new AtomicBoolean(); + var late = new AtomicInteger(); + var client = readyClient(closed, late); + var queue = new LinkedBlockingQueue(); + var workers = Executors.newCachedThreadPool(); + var started = new CountDownLatch(1); + var release = new CountDownLatch(1); + var calls = new AtomicInteger(); + var setup = setup(client, workers, queue, closed); + try { + setup.context.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + calls.incrementAndGet(); + if (state == 1) return WaitForConditionResult.continuePolling(2); + started.countDown(); + await(release); + return WaitForConditionResult.stopPolling(state); + }, + waitConfig()); + var task = queue.poll(3, TimeUnit.SECONDS); + assertNotNull(task); + CompletableFuture.runAsync(task).get(3, TimeUnit.SECONDS); + assertTrue(started.await(3, TimeUnit.SECONDS)); + var closing = CompletableFuture.runAsync(setup.manager::close); + assertTrue(setup.manager.closeEntered.await(3, TimeUnit.SECONDS)); + assertThrows(TimeoutException.class, () -> closing.get(150, TimeUnit.MILLISECONDS)); + release.countDown(); + closing.get(3, TimeUnit.SECONDS); + assertEquals(2, calls.get()); + assertEquals( + OperationStatus.SUCCEEDED, + client.getOperationByName("condition").status()); + var checkpoints = client.getOperationUpdates().size(); + var replay = setup(client, workers, new LinkedBlockingQueue<>(), new AtomicBoolean()); + try { + assertEquals( + 2, + replay.context.waitForCondition( + "condition", + Integer.class, + (state, step) -> { + calls.incrementAndGet(); + throw new AssertionError("A completed predicate must not run on replay"); + }, + waitConfig())); + assertEquals(2, calls.get()); + assertEquals(checkpoints, client.getOperationUpdates().size()); + assertEquals(0, late.get()); + } finally { + replay.manager.close(); + } + setup.manager + .runCheckpointContinuation(() -> fail("New work must be rejected after close")) + .get(3, TimeUnit.SECONDS); + assertTrue(queue.isEmpty()); + } finally { + release.countDown(); + setup.manager.close(); + workers.shutdownNow(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + } + } + + @Test + void earlyParallelCompletionStopsOnlyItsQueuedConditionAndUnblocksTheWaitingBranch() throws Exception { + var closed = new AtomicBoolean(); + var late = new AtomicInteger(); + var client = readyClient(closed, late); + var queue = new LinkedBlockingQueue(); + var workers = Executors.newCachedThreadPool(); + var created = new CountDownLatch(1); + var calls = new AtomicInteger(); + var condition = new AtomicReference(); + var setup = setup(client, workers, queue, closed); + try { + var parallel = setup.context.parallel( + "early", + ParallelConfig.builder() + .completionConfig(CompletionConfig.minSuccessful(1)) + .build()); + var task = new AtomicReference(); + try (parallel) { + parallel.branch("waiting", Integer.class, branch -> { + var future = branch.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + calls.incrementAndGet(); + return WaitForConditionResult.continuePolling(state + 1); + }, + waitConfig()); + condition.set((BaseDurableOperation) future); + created.countDown(); + return future.get(); + }); + parallel.branch("winner", String.class, branch -> { + await(created); + try { + task.set(queue.poll(3, TimeUnit.SECONDS)); + assertNotNull(task.get()); + condition.get().getRunningUserHandler().get(3, TimeUnit.SECONDS); + } catch (Exception failure) { + throw new AssertionError(failure); + } + return "winner"; + }); + } + var result = parallel.get(); + assertEquals(1, result.succeeded()); + assertFalse(condition.get().getCompletionFuture().isDone()); + var checkpointCount = client.getOperationUpdates().size(); + var closing = CompletableFuture.runAsync(setup.manager::close); + assertTrue(setup.manager.closeEntered.await(3, TimeUnit.SECONDS)); + assertThrows(TimeoutException.class, () -> closing.get(150, TimeUnit.MILLISECONDS)); + CompletableFuture.runAsync(task.get()).get(3, TimeUnit.SECONDS); + closing.get(3, TimeUnit.SECONDS); + assertTrue(condition.get().getCompletionFuture().isCompletedExceptionally()); + assertEquals(1, calls.get()); + assertEquals(0, late.get()); + assertEquals(checkpointCount, client.getOperationUpdates().size()); + assertEquals(1, result.succeeded(), "Cleanup must not rewrite the stored early result"); + } finally { + assertThrows(SuspendExecutionException.class, setup.manager::suspendExecution); + Runnable pending; + while ((pending = queue.poll()) != null) pending.run(); + workers.shutdownNow(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + setup.manager.close(); + } + } + + private record Setup(ControlledManager manager, DurableContextImpl context) {} + + private Setup setup( + LocalMemoryExecutionClient client, + ExecutorService workers, + BlockingQueue queue, + AtomicBoolean closed) { + var config = DurableConfig.builder() + .withDurableExecutionClient(client) + .withExecutorService(workers) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO) + .build(); + var manager = new ControlledManager(input(client), config, queue, closed); + rememberContext(manager); + manager.registerActiveThread(null); + manager.setCurrentThreadContext(new ThreadContext(null, ThreadType.CONTEXT)); + var context = DurableContextImpl.createRootContext(manager, config, null); + BaseContextImpl.setCurrentContext(context); + return new Setup(manager, context); + } + + private static WaitForConditionConfig waitConfig() { + return WaitForConditionConfig.builder() + .initialState(1) + .waitStrategy((state, attempt) -> Duration.ofSeconds(1)) + .build(); + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + + private static LocalMemoryExecutionClient readyClient(AtomicBoolean closed, AtomicInteger late) { + return new LocalMemoryExecutionClient() { + public CheckpointDurableExecutionResponse checkpoint( + String arn, String token, List updates) { + if (closed.get()) late.incrementAndGet(); + var applied = super.checkpoint(arn, token, updates); + advanceTime(); + var ready = super.checkpoint(arn, applied.checkpointToken(), List.of()); + var states = new LinkedHashMap(); + applied.newExecutionState().operations().forEach(op -> states.put(op.id(), op)); + ready.newExecutionState().operations().forEach(op -> states.put(op.id(), op)); + return ready.toBuilder() + .newExecutionState(CheckpointUpdatedExecutionState.builder() + .operations(states.values()) + .build()) + .build(); + } + }; + } + + @ParameterizedTest + @ValueSource(strings = {"queued", "cancelled-observation", "after-old-worker-join", "already-suspended"}) + void normalCloseQuiescesAdmittedReadyWork(String mode) throws Exception { + var closed = new AtomicBoolean(); + var lateBackend = new AtomicInteger(); + var predicates = new AtomicInteger(); + var queue = new LinkedBlockingQueue(); + var workers = new GateExecutor(mode.equals("after-old-worker-join")); + var client = new LocalMemoryExecutionClient() { + public CheckpointDurableExecutionResponse checkpoint( + String arn, String token, List updates) { + if (closed.get()) lateBackend.incrementAndGet(); + var applied = super.checkpoint(arn, token, updates); + advanceTime(); + var ready = super.checkpoint(arn, applied.checkpointToken(), List.of()); + var states = new LinkedHashMap(); + applied.newExecutionState().operations().forEach(op -> states.put(op.id(), op)); + ready.newExecutionState().operations().forEach(op -> states.put(op.id(), op)); + return ready.toBuilder() + .newExecutionState(CheckpointUpdatedExecutionState.builder() + .operations(states.values()) + .build()) + .build(); + } + }; + var config = DurableConfig.builder() + .withDurableExecutionClient(client) + .withExecutorService(workers) + .withCheckpointDelay(Duration.ZERO) + .withPollingStrategy(attempt -> Duration.ZERO) + .build(); + var manager = new ControlledManager(input(client), config, queue, closed); + rememberContext(manager); + manager.registerActiveThread(null); + manager.setCurrentThreadContext(new ThreadContext(null, ThreadType.CONTEXT)); + var context = DurableContextImpl.createRootContext(manager, config, null); + BaseContextImpl.setCurrentContext(context); + var body = new CompletableFuture(); + var selected = manager.runUntilCompleteOrSuspend(body); + try { + context.waitForConditionAsync( + "condition", + Integer.class, + (state, step) -> { + predicates.incrementAndGet(); + return state == 1 + ? WaitForConditionResult.continuePolling(2) + : WaitForConditionResult.stopPolling(state); + }, + WaitForConditionConfig.builder() + .initialState(1) + .waitStrategy((state, attempt) -> Duration.ofSeconds(1)) + .build()); + var task = queue.poll(3, TimeUnit.SECONDS); + assertNotNull(task); + assertTrue(workers.firstFinished.await(3, TimeUnit.SECONDS)); + if (mode.equals("cancelled-observation")) assertTrue(manager.observation.cancel(false)); + CompletableFuture taskRun = null; + if (mode.equals("after-old-worker-join")) { + taskRun = CompletableFuture.runAsync(task); + assertTrue(workers.nextDispatchEntered.await(3, TimeUnit.SECONDS)); + } + var stopped = mode.equals("already-suspended"); + if (stopped) assertThrows(SuspendExecutionException.class, manager::suspendExecution); + else { + body.complete("root-done"); + assertEquals("root-done", selected.get(3, TimeUnit.SECONDS)); + } + var closing = CompletableFuture.runAsync(manager::close); + assertTrue(manager.closeEntered.await(3, TimeUnit.SECONDS)); + boolean returnedWhileHeld; + try { + closing.get(150, TimeUnit.MILLISECONDS); + returnedWhileHeld = true; + } catch (TimeoutException expected) { + returnedWhileHeld = false; + } + workers.releaseNextDispatch.countDown(); + if (taskRun == null) taskRun = CompletableFuture.runAsync(task); + taskRun.get(3, TimeUnit.SECONDS); + closing.get(3, TimeUnit.SECONDS); + workers.shutdown(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + var returnedEarly = returnedWhileHeld; + System.out.println("CONTINUATION_CLOSE mode=" + mode + " returnedWhileHeld=" + returnedEarly + + " predicateCalls=" + predicates.get() + " lateCheckpointAttempts=" + + manager.lateCheckpointAttempts.get() + " lateBackendCalls=" + lateBackend.get()); + assertAll( + () -> { + if (!stopped) assertFalse(returnedEarly, "Close must wait for the actual admitted task"); + }, + () -> assertEquals(1, predicates.get(), "No new predicate may run after teardown"), + () -> assertEquals(0, manager.lateCheckpointAttempts.get()), + () -> assertEquals(0, lateBackend.get())); + if (!stopped) assertEquals("root-done", selected.join()); + } finally { + workers.releaseNextDispatch.countDown(); + Runnable pending; + while ((pending = queue.poll()) != null) pending.run(); + workers.shutdownNow(); + assertTrue(workers.awaitTermination(3, TimeUnit.SECONDS)); + manager.close(); + } + } + + private static final class ControlledManager extends ExecutionManager { + private final BlockingQueue queue; + private final AtomicBoolean closed; + private final CountDownLatch closeEntered = new CountDownLatch(1); + private final AtomicInteger lateCheckpointAttempts = new AtomicInteger(); + private volatile CompletableFuture observation; + + private ControlledManager( + DurableExecutionInput input, + DurableConfig config, + BlockingQueue queue, + AtomicBoolean closed) { + super(input, config, null); + this.queue = queue; + this.closed = closed; + } + + public CompletableFuture runCheckpointContinuation(BaseDurableOperation owner, Runnable task) { + observation = super.runCheckpointContinuation(owner, task, queue::add); + return observation; + } + + public CompletableFuture sendOperationUpdate(OperationUpdate update) { + if (closed.get()) lateCheckpointAttempts.incrementAndGet(); + return super.sendOperationUpdate(update); + } + + public void close() { + closeEntered.countDown(); + super.close(); + closed.set(true); + } + } + + private static final class GateExecutor extends ThreadPoolExecutor { + private final AtomicInteger submitted = new AtomicInteger(); + private final AtomicInteger finished = new AtomicInteger(); + private final boolean hold; + private final CountDownLatch firstFinished = new CountDownLatch(1); + private final CountDownLatch nextDispatchEntered = new CountDownLatch(1); + private final CountDownLatch releaseNextDispatch = new CountDownLatch(1); + + private GateExecutor(boolean hold) { + super(1, 1, 0, TimeUnit.SECONDS, new LinkedBlockingQueue<>()); + this.hold = hold; + } + + public void execute(Runnable task) { + if (submitted.incrementAndGet() == 2 && hold) { + nextDispatchEntered.countDown(); + try { + assertTrue(releaseNextDispatch.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + super.execute(task); + } + + protected void afterExecute(Runnable task, Throwable failure) { + if (finished.incrementAndGet() == 1) firstFinished.countDown(); + } + } + + private static DurableExecutionInput input(LocalMemoryExecutionClient client) { + var states = new ArrayList<>(client.getAllOperations()); + states.add( + 0, + Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.EPOCH) + .executionDetails(ExecutionDetails.builder() + .inputPayload("\"input\"") + .build()) + .build()); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/close/execution", + "token", + CheckpointUpdatedExecutionState.builder().operations(states).build()); + } +} diff --git a/sdk/src/main/java/software/amazon/lambda/durable/DurableConfig.java b/sdk/src/main/java/software/amazon/lambda/durable/DurableConfig.java index 54476ec79..ee81fcd23 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/DurableConfig.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/DurableConfig.java @@ -497,8 +497,9 @@ public Builder withCheckpointEmptyMap(boolean checkpointEmptyMap) { * constructed again when its exact concrete type is already explicitly configured. Unrelated plugins retain * their existing registration behavior. * - *

    Calling this method replaces any previously registered plugins. Plugins are called in registration order. - * A fresh builder combines this explicit list with environment-selected plugins. On a builder returned by + *

    Calling this method replaces any previously registered plugins. Hooks use registration order, except + * invocation end, which uses reverse order under the 2.2.2 same-thread lifecycle contract. A fresh builder + * combines this explicit list with environment-selected plugins. On a builder returned by * {@link DurableConfig#toBuilder()}, this method replaces the complete resolved list and dynamic discovery * remains disabled, including when the replacement list is empty. * 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..4315ed4ac 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 @@ -10,7 +10,12 @@ public class UnrecoverableDurableExecutionException extends DurableExecutionExce private final boolean retryable; public UnrecoverableDurableExecutionException(ErrorObject errorObject, boolean retryable) { - super(errorObject.errorMessage()); + this(errorObject, retryable, null); + } + + /** Creates an invocation control failure while retaining its original in-process cause. */ + public UnrecoverableDurableExecutionException(ErrorObject errorObject, boolean retryable, Throwable cause) { + super(errorObject.errorMessage(), cause); this.errorObject = errorObject; this.retryable = retryable; } 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..0430d6ab1 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 @@ -198,6 +198,11 @@ List fetchAllPages(CheckpointUpdatedExecutionState checkpointUpdatedE private void checkpointBatch(List updates) { synchronized (pollingFutures) { + // A READY recheck can complete a poll before its scheduled batch. Keep only consumers still waiting. + pollingFutures.entrySet().removeIf(entry -> { + entry.getValue().removeIf(CompletableFuture::isDone); + return entry.getValue().isEmpty(); + }); // filter the null values from pollers var request = updates.stream().filter(Objects::nonNull).toList(); 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..8b97b9e27 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 @@ -4,13 +4,22 @@ import com.amazonaws.services.lambda.runtime.Context; import com.amazonaws.services.lambda.runtime.RequestHandler; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.UndeclaredThrowableException; import java.nio.charset.StandardCharsets; +import java.util.Collections; +import java.util.IdentityHashMap; import java.util.concurrent.CompletableFuture; import java.util.concurrent.CompletionException; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.Executor; import java.util.concurrent.atomic.AtomicReference; import java.util.function.BiFunction; +import java.util.function.Function; +import java.util.function.Supplier; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import org.slf4j.MDC; import software.amazon.awssdk.services.lambda.model.ErrorObject; import software.amazon.awssdk.services.lambda.model.Operation; import software.amazon.awssdk.services.lambda.model.OperationAction; @@ -26,6 +35,7 @@ import software.amazon.lambda.durable.logging.DurableLogger; import software.amazon.lambda.durable.model.DurableExecutionInput; import software.amazon.lambda.durable.model.DurableExecutionOutput; +import software.amazon.lambda.durable.model.SafeCloseable; import software.amazon.lambda.durable.plugin.InvocationEndInfo; import software.amazon.lambda.durable.plugin.InvocationInfo; import software.amazon.lambda.durable.plugin.InvocationStatus; @@ -49,180 +59,395 @@ public class DurableExecutor { private DurableExecutor() {} + /** + * Returns whether this core runs invocation start/end hooks on the root handler thread and awaits handler cleanup + * and end-hook completion before returning. Plugins that retain thread-local scopes may require this capability + * before construction. Introduced with the 2.2.2 lifecycle contract. + */ + public static boolean supportsSameThreadInvocationHooks() { + return true; + } + public static DurableExecutionOutput execute( DurableExecutionInput input, Context lambdaContext, TypeToken inputType, BiFunction handler, DurableConfig config) { - var pluginRunner = config.getPluginRunner(); - try (var executionManager = new ExecutionManager(input, config, lambdaContext)) { - var isFirstInvocation = !executionManager.isReplaying(); - var requestId = lambdaContext != null ? lambdaContext.getAwsRequestId() : null; - var executionArn = input.durableExecutionArn(); - - 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)); - - // 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); - } - - 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 - - // 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 (var manager = new ExecutionManager(input, config, lambdaContext)) { + manager.registerActiveThread(ROOT_THREAD_ID); + var invocation = new Invocation<>(input, lambdaContext, inputType, handler, config, manager); 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); - - // 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; - } - - // 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)); + return invocation.execute().join(); + } catch (CompletionException failure) { + ExceptionHelper.sneakyThrow(ExceptionHelper.unwrapCompletableFuture(failure)); return null; } } } - private static void fireOnInvocationEnd( - PluginRunner pluginRunner, - ExecutionManager executionManager, - String requestId, - String executionArn, - boolean isFirstInvocation, - InvocationStatus status, - Throwable error, - Object executionInput, - Object executionResult) { - if (pluginRunner.isEmpty()) { - return; + /** Invocation-local state, accessed on the handler thread when plugins are present. */ + private static final class Invocation { + private final Context lambdaContext; + private final TypeToken inputType; + private final BiFunction handler; + private final DurableConfig config; + private final ExecutionManager manager; + private final PluginRunner plugins; + private final String requestId; + private final String executionArn; + private final boolean isFirstInvocation; + private I userInput; + private boolean started; + + private Invocation( + DurableExecutionInput input, + Context lambdaContext, + TypeToken inputType, + BiFunction handler, + DurableConfig config, + ExecutionManager manager) { + this.lambdaContext = lambdaContext; + this.inputType = inputType; + this.handler = handler; + this.config = config; + this.manager = manager; + plugins = config.getPluginRunner(); + requestId = lambdaContext != null ? lambdaContext.getAwsRequestId() : null; + executionArn = input.durableExecutionArn(); + isFirstInvocation = !manager.isReplaying(); + } + + private CompletableFuture execute() { + var body = new CompletableFuture(); + // Select the invocation outcome before scheduling, including with an inline executor. A suspension or + // termination can win while the handler is still unwinding; its finally block must not replace it. + var outcome = manager.runUntilCompleteOrSuspend(body).handle(Outcome::new); + Supplier task = () -> { + Outcome.capture(this::invokeHandler).complete(body); + var selected = outcome.join(); + return finishInvocation(selected.value(), selected.failure()); + }; + // Cleanup is awaited even without plugins. Keep that path free of plugin MDC handling. + if (plugins.isEmpty()) return CompletableFuture.supplyAsync(task, config.getExecutorService()); + return supplyHandler(task, this::finishFailure, config.getExecutorService()); + } + + private O invokeHandler() { + manager.setCurrentThreadContext(new ThreadContext(ROOT_THREAD_ID, ThreadType.CONTEXT)); + Throwable inputFailure = null; + try { + userInput = extractUserInput(manager.getExecutionOperation(), config.getSerDes(), inputType); + } catch (Throwable failure) { + // Deserialize only once. Even failed input gets paired start/end hooks with a null input value. + inputFailure = failure; + } + fireOnInvocationStart(); + if (inputFailure != null) ExceptionHelper.sneakyThrow(inputFailure); + var context = DurableContextImpl.createRootContext(manager, config, lambdaContext); + DurableContextImpl.setCurrentContext(context); + try (var ignored = DurableLogger.attachContext()) { + return handler.apply(userInput, context); + } + } + + private void fireOnInvocationStart() { + if (plugins.isEmpty()) return; + var info = new InvocationInfo( + requestId, + executionArn, + isFirstInvocation, + manager.getExecutionOperation().startTimestamp(), + userInput, + PluginInfoConverter.toOperationItemMap( + manager.getOperationsSnapshot(), manager.getInitialOperationIds()), + PluginInfoConverter.toOperationItemMap( + manager.getUpdatedOperationsSnapshot(), manager.getInitialOperationIds())); + started = true; + plugins.onInvocationStart(info); + } + + private DurableExecutionOutput finishInvocation(O value, Throwable failure) { + if (failure != null) return finishFailure(ExceptionHelper.unwrapCompletableFuture(failure)); + DurableExecutionOutput output; + try { + var payload = config.getSerDes().serialize(value); + output = DurableExecutionOutput.success(handleLargePayload(manager, payload)); + } catch (Throwable deliveryFailure) { + return failDelivery(deliveryFailure); + } + fireOnInvocationEnd(InvocationStatus.SUCCEEDED, null, value); + return output; + } + + private DurableExecutionOutput finishFailure(Throwable cause) { + var status = failureStatus(cause); + if (status == InvocationStatus.FAILED) return finishTerminalFailure(cause); + try { + fireOnInvocationEnd(status, status == InvocationStatus.PENDING ? null : cause, null); + } catch (Error endFailure) { + if (status == InvocationStatus.RETRYING) + ExceptionHelper.sneakyThrow(combinePreparationAndEndFailures(cause, endFailure)); + throw endFailure; + } + if (status == InvocationStatus.PENDING) return DurableExecutionOutput.pending(); + ExceptionHelper.sneakyThrow(cause); + return null; + } + + private DurableExecutionOutput finishTerminalFailure(Throwable cause) { + DurableExecutionOutput output; + try { + output = DurableExecutionOutput.failure(buildErrorObject(cause, config.getSerDes())); + } catch (Throwable deliveryFailure) { + return failDelivery(deliveryFailure); + } + fireOnInvocationEnd(InvocationStatus.FAILED, cause, null); + return output; + } + + private DurableExecutionOutput failDelivery(Throwable deliveryFailure) { + // Serialization/checkpointing did not produce a terminal response. Close plugin resources with RETRYING + // before preserving the original delivery failure for the Lambda caller. + var cause = ExceptionHelper.unwrapCompletableFuture(deliveryFailure); + try { + fireOnInvocationEnd(InvocationStatus.RETRYING, cause, null); + } catch (Error endFailure) { + deliveryFailure = combinePreparationAndEndFailures(cause, endFailure); + } + ExceptionHelper.sneakyThrow(deliveryFailure); + return null; + } + + private void fireOnInvocationEnd(InvocationStatus status, Throwable error, Object result) { + if (!started) return; + plugins.onInvocationEnd(new InvocationEndInfo( + requestId, + executionArn, + isFirstInvocation, + manager.getExecutionOperation().startTimestamp(), + PluginInfoConverter.toOperationItemMap( + manager.getOperationsSnapshot(), manager.getInitialOperationIds()), + status, + error, + userInput, + result)); + } + } + + /** Preserves preparation as primary unless cleanup introduces the first JVM-fatal failure. */ + @SuppressWarnings("removal") + private static Throwable combinePreparationAndEndFailures(Throwable preparation, Error cleanup) { + try { + rethrowLifecycleFatal(preparation); + } catch (VirtualMachineError | ThreadDeath fatal) { + if (fatal != cleanup) fatal.addSuppressed(cleanup); + return fatal; + } + if (cleanup instanceof VirtualMachineError || cleanup instanceof ThreadDeath) { + if (cleanup != preparation) cleanup.addSuppressed(preparation); + return cleanup; + } + if (preparation != cleanup) preparation.addSuppressed(cleanup); + return preparation; + } + + private static InvocationStatus failureStatus(Throwable failure) { + if (failure instanceof SuspendExecutionException) return InvocationStatus.PENDING; + if (failure instanceof UnrecoverableDurableExecutionException unrecoverable && unrecoverable.isRetryable()) { + return InvocationStatus.RETRYING; + } + return InvocationStatus.FAILED; + } + + /** Captures task completion without running lifecycle hooks on CompletableFuture completion threads. */ + private record Outcome(T value, Throwable failure) { + private static Outcome capture(Supplier task) { + try { + return new Outcome<>(task.get(), null); + } catch (Throwable failure) { + return new Outcome<>(null, failure); + } + } + + private void complete(CompletableFuture future) { + if (failure == null) future.complete(value); + else future.completeExceptionally(failure); + } + } + + @SuppressWarnings("removal") + private static CompletableFuture supplyHandler( + Supplier task, Function initializationFailure, Executor executor) { + var result = new CompletableFuture(); + var classifiedWorkerFailure = new AtomicReference(); + Runnable work = (Runnable & CompletableFuture.AsynchronousCompletionTask) () -> { + SafeCloseable restore; + try { + restore = restoreMdcOnClose(); + } catch (Throwable failure) { + Throwable initializationCause; + try { + initializationCause = normalizeMdcInitializationFailure(failure); + } catch (VirtualMachineError | ThreadDeath fatal) { + // Wake the invocation caller with the original fatal before it escapes the actual worker. + // Throwing first would leave the observation future incomplete on an asynchronous executor. + result.completeExceptionally(fatal); + throw fatal; + } + Outcome.capture(() -> initializationFailure.apply(initializationCause)) + .complete(result); + return; + } + var outcome = Outcome.capture(task); + var restoringAfterNonfatalOutcome = false; + try { + try { + // End/delivery failures must also escape their actual worker. Handler-body failures have + // already been mapped to their selected durable outcome; this does not reclassify them. + rethrowLifecycleFatal(outcome.failure()); + } catch (VirtualMachineError | ThreadDeath fatal) { + closeMdcAfterFatal(restore, fatal); + throw fatal; + } + restoringAfterNonfatalOutcome = true; + restore.close(); + } catch (Throwable workerFailure) { + try { + rethrowLifecycleFatal(workerFailure); + clearMdcAfterOrdinaryFailure(workerFailure); + } catch (VirtualMachineError | ThreadDeath fatal) { + if (restoringAfterNonfatalOutcome && outcome.failure() != null && outcome.failure() != fatal) { + fatal.addSuppressed(outcome.failure()); + } + // End describes the already selected SDK outcome, not successful return to the runtime. + // Publish this fatal before throwing it from the worker; do not dispatch End again. + result.completeExceptionally(fatal); + throw fatal; + } + classifiedWorkerFailure.set(workerFailure); + throw workerFailure; + } finally { + // Ordinary restoration failures retain the selected outcome. A fatal has already settled result. + outcome.complete(result); + } + }; + try { + executor.execute(work); + } catch (Throwable dispatchFailure) { + if (!result.isDone()) throw dispatchFailure; + // Inline execution can throw from MDC restoration after settling the selected outcome. Preserve that + // outcome for an ordinary cleanup failure, while a JVM-fatal failure still reaches the caller. + // An inline task already inspected this exact ordinary failure. Do not re-read custom diagnostics; + // a different failure raised by the executor itself still receives the existing classification. + if (dispatchFailure != classifiedWorkerFailure.get()) rethrowLifecycleFatal(dispatchFailure); + } + return result; + } + + @SuppressWarnings("removal") + private static void clearMdcAfterOrdinaryFailure(Throwable restorationFailure) { + try { + MDC.clear(); + } catch (Throwable clearFailure) { + if (clearFailure == restorationFailure) return; // Already classified; do not re-read its diagnostics. + try { + rethrowLifecycleFatal(clearFailure); + } catch (VirtualMachineError | ThreadDeath fatal) { + if (fatal != restorationFailure) fatal.addSuppressed(restorationFailure); + throw fatal; + } + restorationFailure.addSuppressed(clearFailure); } - pluginRunner.onInvocationEnd(new InvocationEndInfo( - requestId, - executionArn, - isFirstInvocation, - executionManager.getExecutionOperation().startTimestamp(), - PluginInfoConverter.toOperationItemMap( - executionManager.getOperationsSnapshot(), executionManager.getInitialOperationIds()), - status, - error, - executionInput, - executionResult)); + } + + /** A repeated fatal object must not be replaced by try-with-resources self-suppression failure. */ + private static void closeMdcAfterFatal(SafeCloseable restore, Error fatal) { + try { + restore.close(); + } catch (Throwable restorationFailure) { + if (restorationFailure == fatal) return; + fatal.addSuppressed(restorationFailure); + try { + rethrowLifecycleFatal(restorationFailure); + // The first fatal still owns caller/worker propagation. Clear only after an ordinary restore failure + // so a ThreadPoolExecutor replacement does not inherit an End hook's invocation-local map. + MDC.clear(); + } catch (Throwable cleanupFailure) { + if (cleanupFailure != fatal && cleanupFailure != restorationFailure) + fatal.addSuppressed(cleanupFailure); + } + } + } + + /** Classifies capture failures once, retaining the ordinary policy of unwrapping only a completion prefix. */ + private static Throwable normalizeMdcInitializationFailure(Throwable failure) { + rethrowDirectMdcFatal(failure); + var visited = Collections.newSetFromMap(new IdentityHashMap()); + var cause = failure; + var normalized = failure; + var completionPrefix = true; + while (visited.add(cause)) { + rethrowDirectMdcFatal(cause); + if (!(cause instanceof CompletionException)) completionPrefix = false; + if (!(cause instanceof CompletionException + || cause instanceof ExecutionException + || cause instanceof InvocationTargetException + || cause instanceof UndeclaredThrowableException)) return normalized; + var next = readMdcCause(cause); + if (next == null || visited.contains(next)) return completionPrefix ? failure : normalized; + if (completionPrefix) normalized = next; + cause = next; + } + return completionPrefix ? failure : normalized; + } + + private static Throwable readMdcCause(Throwable failure) { + try { + return failure.getCause(); + } catch (Throwable unreadableDiagnostic) { + rethrowDirectMdcFatal(unreadableDiagnostic); + return null; + } + } + + @SuppressWarnings("removal") + private static void rethrowDirectMdcFatal(Throwable failure) { + if (failure instanceof VirtualMachineError fatal) throw fatal; + if (failure instanceof ThreadDeath fatal) throw fatal; + } + + /** Inspects only standard transport wrappers at this lifecycle boundary, without trusting diagnostics. */ + @SuppressWarnings("removal") + private static void rethrowLifecycleFatal(Throwable failure) { + if (failure instanceof VirtualMachineError fatal) throw fatal; + if (failure instanceof ThreadDeath fatal) throw fatal; + if (failure == null) return; + var visited = Collections.newSetFromMap(new IdentityHashMap()); + var cause = failure; + while (cause != null && visited.add(cause)) { + if (cause instanceof VirtualMachineError fatal) throw fatal; + if (cause instanceof ThreadDeath fatal) throw fatal; + if (!(cause instanceof CompletionException + || cause instanceof ExecutionException + || cause instanceof InvocationTargetException + || cause instanceof UndeclaredThrowableException)) return; + try { + cause = cause.getCause(); + } catch (VirtualMachineError | ThreadDeath fatal) { + throw fatal; + } catch (Throwable unreadableDiagnostic) { + return; + } + } + } + + private static SafeCloseable restoreMdcOnClose() { + var previous = MDC.getCopyOfContextMap(); + return () -> { + if (previous == null) MDC.clear(); + else MDC.setContextMap(previous); + }; } private static String handleLargePayload(ExecutionManager executionManager, String outputPayload) { 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 4c7feeda4..1aa148c2a 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 @@ -7,6 +7,7 @@ import java.util.ArrayList; import java.util.Collection; import java.util.Collections; +import java.util.HashMap; import java.util.HashSet; import java.util.List; import java.util.Map; @@ -15,11 +16,13 @@ import java.util.concurrent.CancellationException; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.Executor; import java.util.concurrent.ThreadPoolExecutor; import java.util.concurrent.atomic.AtomicReference; 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; @@ -75,8 +78,13 @@ public class ExecutionManager implements SafeCloseable { private final Set activeThreads = Collections.synchronizedSet(new HashSet<>()); private static final ThreadLocal currentThreadContext = new ThreadLocal<>(); private final CompletableFuture executionExceptionFuture = new CompletableFuture<>(); - // Guarded by activeThreads so starting a checkpoint request is atomic with the last-thread suspension decision. + // Requests and polling continuations; guarded by activeThreads so admission is atomic with last-thread suspension. private int checkpointRequestsInFlight; + // Guarded by activeThreads; an observation future can be cancelled independently of its actual task. + private final Map checkpointContinuations = new HashMap<>(); + private boolean closing; + // Guarded by activeThreads; selected against normal close before publishing callbacks outside the monitor. + private UnrecoverableDurableExecutionException selectedContinuationFailure; /** * Per-wait state used to coordinate the caller thread with the future completion callback. @@ -544,8 +552,135 @@ void finishCheckpointProcessing() { } } + /** + * Runs a polling continuation away from the serialized checkpoint batcher. Keeps checkpoint processing active from + * receiving the callback until the continuation has registered its next worker (or finished inline). Canceling the + * returned future affects observation only; queued or running work still owns its activity lease. + * + * @param continuation operation work to dispatch after receiving a polling update + * @return completion of the continuation, or an already-completed future when execution has stopped + */ + public CompletableFuture runCheckpointContinuation(Runnable continuation) { + return runCheckpointContinuation(null, continuation, InternalExecutor.INSTANCE); + } + + /** Schedules an operation continuation whose unfinished waiters are stopped during manager cleanup. */ + public CompletableFuture runCheckpointContinuation(BaseDurableOperation owner, Runnable continuation) { + return runCheckpointContinuation(owner, continuation, InternalExecutor.INSTANCE); + } + + CompletableFuture runCheckpointContinuation(Runnable continuation, Executor coordinator) { + return runCheckpointContinuation(null, continuation, coordinator); + } + + CompletableFuture runCheckpointContinuation( + BaseDurableOperation owner, Runnable continuation, Executor coordinator) { + var registration = registerCheckpointContinuation(owner); + if (registration == null) { + if (isClosing()) stopContinuationOwner(owner); + return CompletableFuture.completedFuture(null); + } + try { + var completion = new CompletableFuture(); + coordinator.execute(() -> completeCheckpointContinuation(owner, registration, continuation, completion)); + return completion; + } catch (RuntimeException | Error failure) { + try { + if ((owner != null && !isClosing()) + || failure instanceof VirtualMachineError + || failure instanceof ThreadDeath) signalContinuationFailure(failure); + } finally { + finishCheckpointContinuation(registration); + } + throw failure; + } + } + + @SuppressWarnings("removal") + private void completeCheckpointContinuation( + BaseDurableOperation owner, + Object registration, + Runnable continuation, + CompletableFuture completion) { + try { + try { + if (!isClosing()) continuation.run(); + } catch (Throwable failure) { + // An operation polling chain can ignore observation. Select retry control before the final lease + // could select PENDING; unowned ordinary helper failures retain their observation-only contract. + if ((owner != null && !isClosing()) + || failure instanceof VirtualMachineError + || failure instanceof ThreadDeath) signalContinuationFailure(failure); + throw failure; + } finally { + finishCheckpointContinuation(registration); + } + completion.complete(null); + } catch (Throwable failure) { + completion.completeExceptionally(failure); + // The observation future may be ignored by a polling chain. Never hide a JVM fatal from its worker. + if (failure instanceof VirtualMachineError fatal) throw fatal; + if (failure instanceof ThreadDeath fatal) throw fatal; + } + } + + private void signalContinuationFailure(Throwable failure) { + var control = new UnrecoverableDurableExecutionException( + ErrorObject.builder() + .errorType(failure.getClass().getName()) + .errorMessage("Error in SDK checkpoint continuation") + .build(), + true, + failure); + UnrecoverableDurableExecutionException selected; + synchronized (activeThreads) { + if (closing && !(failure instanceof VirtualMachineError || failure instanceof ThreadDeath)) return; + if (selectedContinuationFailure == null) selectedContinuationFailure = control; + selected = selectedContinuationFailure; + } + // A later admitted publisher may drain the winner's callbacks. Keep that CompletableFuture behavior, + // but publish/wake using the selected cause, outside the monitor needed by worker cleanup. + if (executionExceptionFuture.completeExceptionally(selected)) stopAllOperations(selected); + } + + private Object registerCheckpointContinuation(BaseDurableOperation owner) { + synchronized (activeThreads) { + if (closing || executionExceptionFuture.isDone()) return null; + var registration = new Object(); + checkpointContinuations.put(registration, owner); + checkpointRequestsInFlight++; + return registration; + } + } + + private void finishCheckpointContinuation(Object registration) { + synchronized (activeThreads) { + // An inline coordinator may rethrow after its task already released this registration. + if (!checkpointContinuations.containsKey(registration)) return; + checkpointContinuations.remove(registration); + try { + finishCheckpointProcessing(); + } finally { + activeThreads.notifyAll(); + } + } + } + + private boolean isClosing() { + synchronized (activeThreads) { + return closing; + } + } + + private static void stopContinuationOwner(BaseDurableOperation owner) { + if (owner != null) owner.getCompletionFuture().completeExceptionally(new SuspendExecutionException()); + } + private boolean shouldSuspendExecution() { - return activeThreads.isEmpty() && checkpointRequestsInFlight == 0 && !executionExceptionFuture.isDone(); + return !closing + && activeThreads.isEmpty() + && checkpointRequestsInFlight == 0 + && !executionExceptionFuture.isDone(); } private void preSuspendCheck() { @@ -595,11 +730,34 @@ public CompletableFuture pollForOperationUpdates(String operationId, /** Shutdown the checkpoint batcher. */ @Override public void close() { + stopCheckpointContinuations(); validateRunningThreads(); checkpointManager.shutdown(); } + private void stopCheckpointContinuations() { + List owners; + synchronized (activeThreads) { + closing = true; + owners = new ArrayList<>(checkpointContinuations.values()); + } + owners.forEach(ExecutionManager::stopContinuationOwner); + synchronized (activeThreads) { + while (!checkpointContinuations.isEmpty()) { + try { + // wait releases this monitor so tasks can finish and release their registrations. No user work + // or future join is performed while holding the coordination monitor. + activeThreads.wait(); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new IllegalStateException( + "Interrupted while waiting for checkpoint continuations", interrupted); + } + } + } + } + private void validateRunningThreads() { // This will detect stuck user thread and thread leaks in the thread pool for (BaseDurableOperation op : registeredOperations.values()) { @@ -645,8 +803,9 @@ public static boolean isTerminalStatus(OperationStatus status) { * @param exception the unrecoverable exception that caused termination */ public void terminateExecution(UnrecoverableDurableExecutionException exception) { - stopAllOperations(exception); + // Select control flow before waking a handler whose finally block can complete its body future. executionExceptionFuture.completeExceptionally(exception); + stopAllOperations(exception); throw exception; } @@ -657,8 +816,8 @@ public void suspendExecution() { private SuspendExecutionException signalSuspension() { var ex = new SuspendExecutionException(); - stopAllOperations(ex); executionExceptionFuture.completeExceptionally(ex); + stopAllOperations(ex); return ex; } 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 a1514e6d5..bdb7fd11a 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 @@ -321,8 +321,21 @@ 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())); + try { + runningUserHandler.set(CompletableFuture.runAsync( + wrapped, getContext().getDurableConfig().getExecutorService())); + } catch (RuntimeException | Error dispatchFailure) { + // Admission failed, so no worker will execute its deregistration finally block. The submitting + // coordinator need not have a user ThreadContext; undo the registration directly on the manager. + if (operationId != null) { + try { + executionManager.deregisterActiveThread(operationId); + } catch (SuspendExecutionException ignored) { + // Preserve the dispatch failure; any suspension was already signaled by deregistration. + } + } + throw dispatchFailure; + } } /** diff --git a/sdk/src/main/java/software/amazon/lambda/durable/operation/StepOperation.java b/sdk/src/main/java/software/amazon/lambda/durable/operation/StepOperation.java index 467a87b94..19fafbdc2 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/operation/StepOperation.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/operation/StepOperation.java @@ -70,7 +70,7 @@ protected void replay(Operation existing) { case STARTED -> { if (isAtMostOnce()) { // AT_MOST_ONCE: treat as interrupted, go through retry logic - handleStepFailure(new StepInterruptedException(existing), attempt); + handleStepFailure(new StepInterruptedException(existing), attempt, null); } else { // AT_LEAST_ONCE: re-execute the step executeStepLogic(attempt); @@ -79,7 +79,7 @@ protected void replay(Operation existing) { // Step is pending retry - Start polling for PENDING -> READY transition case PENDING -> { if (existing.stepDetails() != null && existing.stepDetails().nextAttemptTimestamp() != null) { - pollReadyAndExecuteStepLogic(existing.stepDetails().nextAttemptTimestamp(), attempt); + pollReadyAndExecuteStepLogic(existing.stepDetails().nextAttemptTimestamp(), attempt, null); } else { throw terminateExecutionWithIllegalDurableOperationException( "Unexpected PENDING step without nextAttemptTimestamp: " + getOperationId()); @@ -93,16 +93,24 @@ protected void replay(Operation existing) { } } - private void pollReadyAndExecuteStepLogic(Instant nextAttemptInstant, int attempt) { + private void pollReadyAndExecuteStepLogic( + Instant nextAttemptInstant, int attempt, CompletableFuture> previousWorker) { pollForOperationUpdates(nextAttemptInstant) .thenCompose(op -> op.status() == OperationStatus.READY ? CompletableFuture.completedFuture(op) : pollForOperationUpdates(nextAttemptInstant)) - .thenRun(() -> executeStepLogic(attempt)); + .thenCompose(ignored -> executionManager.runCheckpointContinuation(this, () -> { + // Leave the serialized checkpoint callback before dispatching a configured executor. The lease + // spans retirement of the old attempt and registration of its replacement. + if (previousWorker != null) previousWorker.join().join(); + if (!isOperationCompleted()) executeStepLogic(attempt); + })); } private void executeStepLogic(int attempt) { + var publishedWorker = new CompletableFuture>(); Runnable userHandler = () -> { + if (isOperationCompleted()) return; // use a try-with-resources to // - add thread id/type to thread local when the step starts // - clear logger properties when the step finishes @@ -119,13 +127,14 @@ private void executeStepLogic(int attempt) { handleStepSucceeded(result); } catch (Throwable e) { - handleStepFailure(e, attempt); + handleStepFailure(e, attempt, publishedWorker); } } }; // Execute user provided step code in user-configured executor runUserHandler(userHandler, ThreadType.STEP); + publishedWorker.complete(getRunningUserHandler()); } private void checkpointStarted() { @@ -156,7 +165,8 @@ private void handleStepSucceeded(T result) { sendOperationUpdate(successUpdate); } - private void handleStepFailure(Throwable exception, int attempt) { + private void handleStepFailure( + Throwable exception, int attempt, CompletableFuture> publishedWorker) { exception = ExceptionHelper.unwrapCompletableFuture(exception); if (exception instanceof SuspendExecutionException suspendExecutionException) { throw suspendExecutionException; @@ -189,7 +199,7 @@ private void handleStepFailure(Throwable exception, int attempt) { sendOperationUpdate(retryUpdate); // Poll for READY status and then execute the step again - pollReadyAndExecuteStepLogic(Instant.now().plusSeconds(retryDelayInSeconds), attempt + 1); + pollReadyAndExecuteStepLogic(Instant.now().plusSeconds(retryDelayInSeconds), attempt + 1, publishedWorker); } else { // Send FAIL - retries exhausted var failUpdate = diff --git a/sdk/src/main/java/software/amazon/lambda/durable/operation/WaitForConditionOperation.java b/sdk/src/main/java/software/amazon/lambda/durable/operation/WaitForConditionOperation.java index 72eb65b8e..40eb517ec 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/operation/WaitForConditionOperation.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/operation/WaitForConditionOperation.java @@ -18,6 +18,7 @@ import software.amazon.lambda.durable.exception.DurableOperationException; import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; import software.amazon.lambda.durable.exception.WaitForConditionFailedException; +import software.amazon.lambda.durable.execution.ExecutionManager; import software.amazon.lambda.durable.execution.SuspendExecutionException; import software.amazon.lambda.durable.execution.ThreadType; import software.amazon.lambda.durable.logging.DurableLogger; @@ -63,7 +64,7 @@ protected void start() { @Override protected void replay(Operation existing) { switch (existing.status()) { - case SUCCEEDED, FAILED -> markAlreadyCompleted(); // Check if already completed / failed + case SUCCEEDED, FAILED, CANCELLED, TIMED_OUT, STOPPED -> markAlreadyCompleted(); case PENDING -> pollReadyAndResumeCheckLoop(existing); // Check if pending retry case STARTED, READY -> resumeCheckLoop(existing); default -> @@ -81,7 +82,8 @@ public T get() { var result = (stepDetails != null) ? stepDetails.result() : null; return deserializeResult(result); } else { - var errorObject = op.stepDetails().error(); + var stepDetails = op.stepDetails(); + var errorObject = stepDetails != null ? stepDetails.error() : null; // Attempt to reconstruct and throw the original exception Throwable original = deserializeException(errorObject); @@ -108,68 +110,88 @@ private void resumeCheckLoop(Operation existing) { } private CompletableFuture pollReadyAndResumeCheckLoop(Operation existing) { - return pollForOperationUpdates() - .thenCompose(op -> op.status() == OperationStatus.READY - ? CompletableFuture.completedFuture(op) - : pollForOperationUpdates()) - .thenAccept(this::resumeCheckLoop); + return pollUntilReady() + .thenCompose(op -> executionManager.runCheckpointContinuation(this, () -> { + if (!isOperationCompleted() && op.status() == OperationStatus.READY) resumeCheckLoop(op); + })); + } + + private CompletableFuture pollUntilReady() { + var known = getOperation(); + if (isReadyOrTerminal(known)) return CompletableFuture.completedFuture(known); + var update = pollForOperationUpdates(); + // Register before re-reading: another checkpoint may already have delivered READY before this poll existed. + known = getOperation(); + if (isReadyOrTerminal(known)) { + update.complete(known); + return CompletableFuture.completedFuture(known); + } + return update.thenCompose( + op -> isReadyOrTerminal(op) ? CompletableFuture.completedFuture(op) : pollUntilReady()); + } + + private boolean isReadyOrTerminal(Operation operation) { + if (operation == null) return false; + var status = operation.status(); + if (status == null || status == OperationStatus.UNKNOWN_TO_SDK_VERSION) { + throw terminateExecutionWithIllegalDurableOperationException( + "Unexpected waitForCondition status: " + operation.statusAsString()); + } + return status == OperationStatus.READY || ExecutionManager.isTerminalStatus(status); } private void executeCheckLogic(T currentState, int attempt) { - Runnable userHandler = () -> { + var publishedWorker = new CompletableFuture>(); + runUserHandler(() -> runCheckLoop(currentState, attempt, publishedWorker), ThreadType.STEP); + publishedWorker.complete(getRunningUserHandler()); + } + + private void runCheckLoop(T currentState, int attempt, CompletableFuture> publishedWorker) { + if (!isOperationCompleted()) { var stepContext = getContext().createStepContext(getOperationId(), getName(), attempt); BaseContextImpl.setCurrentContext(stepContext); try (var ignored = DurableLogger.attachContext()) { try { - // Checkpoint START if not already started var existing = getOperation(); if (existing == null || existing.status() != OperationStatus.STARTED) { - var startUpdate = OperationUpdate.builder().action(OperationAction.START); - sendOperationUpdateAsync(startUpdate); + sendOperationUpdateAsync(OperationUpdate.builder().action(OperationAction.START)); } - - // Execute check function inside the plugin hook boundary so a failure is reported - // through onUserFunctionEnd; checkpoint/poll handling stays outside the boundary. - WaitForConditionResult result = - runUserFunction(attempt, () -> checkFunc.apply(currentState, stepContext)); - - // Normalize the value through SerDes so first execution matches replay. + var stateForCheck = currentState; + var result = runUserFunction(attempt, () -> checkFunc.apply(stateForCheck, stepContext)); var serializedState = serializeAndDeserializeResult(result.value()); - T deserializedValue = serializedState.deserialized(); - + var deserializedValue = serializedState.deserialized(); if (result.isDone()) { - // Condition met — checkpoint SUCCEED - var successUpdate = OperationUpdate.builder() + sendOperationUpdate(OperationUpdate.builder() .action(OperationAction.SUCCEED) - .payload(serializedState.serialized()); - sendOperationUpdate(successUpdate); - } else { - // Compute delay from strategy - Duration delay = config.waitStrategy().evaluate(deserializedValue, attempt); - - // Checkpoint RETRY with delay - var retryUpdate = OperationUpdate.builder() - .action(OperationAction.RETRY) - .payload(serializedState.serialized()) - .stepOptions(StepOptions.builder() - .nextAttemptDelaySeconds(Math.toIntExact(delay.toSeconds())) - .build()); - sendOperationUpdate(retryUpdate); - - // Poll for READY, then continue the loop - pollForOperationUpdates() - .thenCompose(op -> op.status() == OperationStatus.READY - ? CompletableFuture.completedFuture(op) - : pollForOperationUpdates()) - .thenRun(() -> executeCheckLogic(deserializedValue, attempt + 1)); + .payload(serializedState.serialized())); + return; } - } catch (Throwable e) { - handleCheckFailure(e); + Duration delay = config.waitStrategy().evaluate(deserializedValue, attempt); + sendOperationUpdate(OperationUpdate.builder() + .action(OperationAction.RETRY) + .payload(serializedState.serialized()) + .stepOptions(StepOptions.builder() + .nextAttemptDelaySeconds(Math.toIntExact(delay.toSeconds())) + .build())); + // Retire even when READY is already available, so a queued sibling can run on a bounded executor. + // The coordinator owns activity until this worker finishes and the next attempt is registered. + continueAfterCurrentWorker(publishedWorker, deserializedValue, attempt + 1); + } catch (Throwable failure) { + handleCheckFailure(failure); + return; } } - }; + } + } - runUserHandler(userHandler, ThreadType.STEP); + private void continueAfterCurrentWorker( + CompletableFuture> publishedWorker, T nextState, int nextAttempt) { + pollUntilReady() + .thenCompose(op -> executionManager.runCheckpointContinuation(this, () -> { + publishedWorker.join().join(); + if (!isOperationCompleted() && op.status() == OperationStatus.READY) + executeCheckLogic(nextState, nextAttempt); + })); } private void handleCheckFailure(Throwable exception) { diff --git a/sdk/src/main/java/software/amazon/lambda/durable/plugin/DurableExecutionPlugin.java b/sdk/src/main/java/software/amazon/lambda/durable/plugin/DurableExecutionPlugin.java index e2f8a46df..d305a1c54 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/plugin/DurableExecutionPlugin.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/plugin/DurableExecutionPlugin.java @@ -18,7 +18,9 @@ public interface DurableExecutionPlugin { // ─── Invocation-level hooks ────────────────────────────────────────── /** - * Called at the start of each Lambda invocation. Use to set up per-invocation state (trace ID, invocation span). + * Called on the root handler thread at the start of each Lambda invocation. Use to set up per-invocation state + * (trace ID, invocation span) and thread-local context. The paired {@link #onInvocationEnd} runs on this same + * thread. * *

    Check {@link InvocationInfo#isFirstInvocation()} to detect the first invocation of an execution (useful for * sampling decisions or execution-level span creation). @@ -26,10 +28,21 @@ public interface DurableExecutionPlugin { default void onInvocationStart(InvocationInfo info) {} /** - * Called at the end of each Lambda invocation. Use to flush spans/metrics before Lambda freezes. + * Called on the same root handler thread as {@link #onInvocationStart}, after the handler and its finally blocks + * have exited. Use to restore thread-local context and flush spans/metrics before Lambda freezes. * *

    This hook is awaited — the SDK blocks until it returns. This is the only safe flush point before Lambda - * freezes the execution environment. + * freezes the execution environment. Suspension and termination also wait for the handler to unwind and this hook + * to finish; a blocked handler, finally block, or end hook delays the invocation response. No caller-thread + * fallback or cleanup timeout is applied. An earlier suspension/termination retains its selected outcome even if a + * later handler finally block returns or throws. + * + *

    Start hooks run in registration order; end hooks run in reverse order to unwind nested thread-local scopes. If + * an end hook throws an Error that is not isolated, the remaining end hooks still run before it is rethrown. The + * first such Error propagates, except that a later VirtualMachineError or ThreadDeath takes precedence over a + * non-JVM-fatal Error. Other distinct end-hook Errors are retained as suppressed failures. If a start hook throws + * an Error that is not isolated, later start hooks are not called, but all configured end hooks still run in + * reverse order. An end hook must therefore tolerate a start hook that did not complete. * *

    Check {@link InvocationEndInfo#invocationStatus()} to detect if the execution reached a terminal state in this * invocation (useful for writing summary records or flushing final data). diff --git a/sdk/src/main/java/software/amazon/lambda/durable/plugin/PluginRunner.java b/sdk/src/main/java/software/amazon/lambda/durable/plugin/PluginRunner.java index ced3638cb..5532eaeb0 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/plugin/PluginRunner.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/plugin/PluginRunner.java @@ -14,8 +14,9 @@ /** * Composes multiple {@link DurableExecutionPlugin} instances into a single dispatcher. * - *

    Event hooks call each plugin in order. Exceptions and nonfatal linkage failures are isolated; other errors, - * including fatal JVM failures, retain their existing propagation behavior. + *

    Event hooks call each plugin in registration order, except invocation end, which unwinds in reverse order. + * Exceptions and nonfatal linkage failures are isolated; other errors propagate. Invocation end finishes the remaining + * cleanup hooks before propagating the first error. * *

    {@code onInvocationEnd} is awaited (the SDK blocks until it returns) to allow plugins to flush data before Lambda * freezes. @@ -82,14 +83,16 @@ public List getPlugins() { /** Calls a void hook on all plugins, isolating exceptions and incompatible binary dependencies. */ private void run(Consumer hook) { - for (var plugin : plugins) { - try { - hook.accept(plugin); - } catch (Exception e) { - logger.warn("Plugin hook threw exception", e); - } catch (LinkageError e) { - logger.warn("Plugin hook could not link a dependency; check SDK/plugin dependency compatibility", e); - } + for (var plugin : plugins) runHook(plugin, hook); + } + + private void runHook(DurableExecutionPlugin plugin, Consumer hook) { + try { + hook.accept(plugin); + } catch (Exception e) { + logger.warn("Plugin hook threw exception", e); + } catch (LinkageError e) { + logger.warn("Plugin hook could not link a dependency; check SDK/plugin dependency compatibility", e); } } @@ -98,11 +101,35 @@ public void onInvocationStart(InvocationInfo info) { } /** - * Called at the end of each invocation. Awaited — the SDK blocks until all plugins return, allowing plugins to - * flush spans/metrics before Lambda freezes. + * Called in reverse registration order on the root handler thread after it unwinds. Awaited — the SDK blocks until + * all plugins return, including during suspension or termination, allowing plugins to flush spans/metrics before + * Lambda freezes. */ public void onInvocationEnd(InvocationEndInfo info) { - run(p -> p.onInvocationEnd(info)); + Error firstError = null; + for (var index = plugins.size() - 1; index >= 0; index--) { + try { + runHook(plugins.get(index), p -> p.onInvocationEnd(info)); + } catch (Error failure) { + // Finish unwinding earlier plugins' thread-local scopes before propagating an end-hook error. + if (firstError == null) { + firstError = failure; + } else if (firstError != failure) { + if (isJvmFatal(failure) && !isJvmFatal(firstError)) { + failure.addSuppressed(firstError); + firstError = failure; + } else { + firstError.addSuppressed(failure); + } + } + } + } + if (firstError != null) throw firstError; + } + + @SuppressWarnings("removal") + private static boolean isJvmFatal(Error failure) { + return failure instanceof VirtualMachineError || failure instanceof ThreadDeath; } public void onOperationStart(OperationInfo info) { diff --git a/sdk/src/test/java/software/amazon/lambda/durable/ReplayValidationTest.java b/sdk/src/test/java/software/amazon/lambda/durable/ReplayValidationTest.java index 98f505ad6..bf62d52fa 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/ReplayValidationTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/ReplayValidationTest.java @@ -3,6 +3,7 @@ package software.amazon.lambda.durable; import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; +import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; @@ -48,7 +49,9 @@ private DurableContext createTestContext(List initialOperations) { null); var context = DurableContextImpl.createRootContext( executionManager, DurableConfig.builder().build(), null); - executionManager.setCurrentThreadContext(new ThreadContext(EXECUTION_OP_ID + "-execution", ThreadType.CONTEXT)); + // Mirror DurableExecutor: the caller must stay active if a worker finishes before step().get() waits. + executionManager.registerActiveThread(null); + executionManager.setCurrentThreadContext(new ThreadContext(null, ThreadType.CONTEXT)); return context; } @@ -59,7 +62,7 @@ void shouldPassValidationWhenNoCheckpointExists() { var context = createTestContext(List.of()); // When & Then: Should not throw - assertDoesNotThrow(() -> context.step("test", String.class, stepCtx -> "result")); + assertEquals("result", context.step("test", String.class, stepCtx -> "result")); } @Test 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..44f0078a9 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/TestUtils.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/TestUtils.java @@ -5,7 +5,7 @@ import static org.mockito.ArgumentMatchers.*; import static org.mockito.Mockito.*; -import java.util.ArrayList; +import java.util.LinkedHashMap; import java.util.List; import java.util.UUID; import software.amazon.awssdk.services.lambda.model.*; @@ -18,7 +18,8 @@ public static DurableExecutionClient createMockClient() { var client = mock(DurableExecutionClient.class); when(client.checkpoint(any(), any(), any())).thenAnswer(invocation -> { var updates = (List) invocation.getArgument(2); - var responseOperations = new ArrayList(); + // A checkpoint returns the latest state of each operation, not one entry per update. + var responseOperations = new LinkedHashMap(); if (updates != null) { for (var update : updates) { @@ -52,14 +53,14 @@ public static DurableExecutionClient createMockClient() { opBuilder.contextDetails(contexDetail.build()); } } - responseOperations.add(opBuilder.build()); + responseOperations.put(update.id(), opBuilder.build()); } } return CheckpointDurableExecutionResponse.builder() .checkpointToken("new-token") .newExecutionState(CheckpointUpdatedExecutionState.builder() - .operations(responseOperations) + .operations(responseOperations.values()) .build()) .build(); }); diff --git a/sdk/src/test/java/software/amazon/lambda/durable/TestUtilsTest.java b/sdk/src/test/java/software/amazon/lambda/durable/TestUtilsTest.java new file mode 100644 index 000000000..1a2864de9 --- /dev/null +++ b/sdk/src/test/java/software/amazon/lambda/durable/TestUtilsTest.java @@ -0,0 +1,51 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable; + +import static org.junit.jupiter.api.Assertions.*; + +import java.util.List; +import java.util.Set; +import org.junit.jupiter.api.Test; +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.plugin.PluginInfoConverter; + +class TestUtilsTest { + @Test + void batchedUpdatesReturnLatestStatePerOperation() { + var client = TestUtils.createMockClient(); + var response = client.checkpoint( + "arn:execution", + "token", + List.of( + update("first", OperationAction.START, null), + update("second", OperationAction.START, null), + update("first", OperationAction.SUCCEED, "\"result\""))); + var operations = response.newExecutionState().operations(); + + // Duplicate IDs make the real SDK callback reject a successful checkpoint. + var change = assertDoesNotThrow(() -> PluginInfoConverter.toOperationChangeInfo( + "request", "arn:execution", operations, operations, Set.of())); + assertEquals(2, operations.size()); + assertEquals(Set.of("first", "second"), change.updatedOperations().keySet()); + assertEquals( + OperationStatus.SUCCEEDED, + change.updatedOperations().get("first").status()); + assertEquals("\"result\"", change.updatedOperations().get("first").result()); + assertEquals( + OperationStatus.STARTED, + change.updatedOperations().get("second").status()); + } + + private static OperationUpdate update(String id, OperationAction action, String payload) { + return OperationUpdate.builder() + .id(id) + .type(OperationType.STEP) + .action(action) + .payload(payload) + .build(); + } +} diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/CompletedPollBatchTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/CompletedPollBatchTest.java new file mode 100644 index 000000000..0646aaa46 --- /dev/null +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/CompletedPollBatchTest.java @@ -0,0 +1,99 @@ +// 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.ArgumentMatchers.*; +import static org.mockito.Mockito.*; + +import java.time.Duration; +import java.util.List; +import java.util.concurrent.*; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.client.DurableExecutionClient; + +class CompletedPollBatchTest { + @ParameterizedTest + @ValueSource( + strings = {"ready", "terminal", "exceptional", "cancelled", "checkpoint", "same-id", "other-id", "pending"}) + void completedPollsDoNotCauseEmptyRpcButLivePollsAndUpdatesStillFlush(String mode) throws Exception { + var client = mock(DurableExecutionClient.class); + var starts = new AtomicInteger(); + var finishes = new AtomicInteger(); + var callbacks = new AtomicInteger(); + var config = DurableConfig.builder() + .withDurableExecutionClient(client) + .withCheckpointDelay(Duration.ofHours(1)) + .build(); + var manager = new CheckpointManager( + config, + "arn:test", + "token", + values -> callbacks.incrementAndGet(), + () -> { + starts.incrementAndGet(); + return true; + }, + finishes::incrementAndGet); + var ready = Operation.builder() + .id("watched") + .type(OperationType.STEP) + .status(OperationStatus.READY) + .build(); + var terminal = ready.toBuilder().status(OperationStatus.SUCCEEDED).build(); + var liveId = mode.equals("other-id") ? "other" : "watched"; + var response = ready.toBuilder().id(liveId).build(); + when(client.checkpoint(anyString(), anyString(), anyList())) + .thenReturn(CheckpointDurableExecutionResponse.builder() + .checkpointToken("next") + .newExecutionState(CheckpointUpdatedExecutionState.builder() + .operations(response) + .build()) + .build()); + var update = OperationUpdate.builder() + .id("write") + .type(OperationType.STEP) + .action(OperationAction.START) + .build(); + try { + var poll = manager.pollForUpdate("watched", attempt -> Duration.ofHours(1)); + CompletableFuture live = null; + if (mode.equals("same-id") || mode.equals("other-id")) + live = manager.pollForUpdate(liveId, attempt -> Duration.ofHours(1)); + switch (mode) { + case "exceptional" -> poll.completeExceptionally(new IllegalStateException("already failed")); + case "cancelled" -> poll.cancel(false); + case "terminal" -> poll.complete(terminal); + case "pending" -> {} + default -> poll.complete(ready); + } + var checkpoint = mode.equals("checkpoint") ? manager.checkpoint(update) : null; + // Force the real delayed batch now, without closing CheckpointManager (which would clear the poll map). + // This exercises the same pending batch as its timer, without a timing-based no-extra-call assertion. + var field = CheckpointManager.class.getDeclaredField("checkpointApiRequestDelayedBatcher"); + field.setAccessible(true); + var dispatcher = (ApiRequestDelayedBatcher) field.get(manager); + dispatcher.shutdown(); + boolean required = live != null || checkpoint != null || mode.equals("pending"); + System.out.println("COMPLETED_POLL mode=" + mode + " APIleases=" + starts.get() + " finished=" + + finishes.get() + " callbacks=" + callbacks.get()); + verify(client, times(required ? 1 : 0)) + .checkpoint(eq("arn:test"), eq("token"), eq(checkpoint != null ? List.of(update) : List.of())); + assertEquals(required ? 1 : 0, starts.get()); + assertEquals(starts.get(), finishes.get()); + assertEquals(required ? 1 : 0, callbacks.get()); + if (live != null) assertEquals(response, live.get(3, TimeUnit.SECONDS)); + if (checkpoint != null) checkpoint.get(3, TimeUnit.SECONDS); + if (mode.equals("pending")) assertEquals(response, poll.get(3, TimeUnit.SECONDS)); + else if (mode.equals("cancelled")) assertTrue(poll.isCancelled()); + else if (mode.equals("exceptional")) assertTrue(poll.isCompletedExceptionally()); + else assertSame(mode.equals("terminal") ? terminal : ready, poll.getNow(null)); + } finally { + manager.shutdown(); + } + } +} diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/ContinuationFailureCloseRaceTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/ContinuationFailureCloseRaceTest.java new file mode 100644 index 000000000..2f80d921a --- /dev/null +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/ContinuationFailureCloseRaceTest.java @@ -0,0 +1,133 @@ +// 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 java.util.concurrent.*; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import software.amazon.awssdk.services.lambda.model.*; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TestUtils; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.operation.BaseDurableOperation; + +class ContinuationFailureCloseRaceTest { + @ParameterizedTest + @CsvSource({"body,true", "reject,true", "body,false", "reject,false"}) + void failureSelectionIsOrderedWithCloseAndDoesNotRunCallbacksUnderCoordinationLock(String mode, boolean closeFirst) + throws Exception { + var input = new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/close-race/execution", + "token", + CheckpointUpdatedExecutionState.builder() + .operations(Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .build()) + .build()); + var manager = new ExecutionManager( + input, + DurableConfig.builder() + .withDurableExecutionClient(TestUtils.createMockClient()) + .build(), + null); + manager.registerActiveThread("root"); + var callbacks = new CountDownLatch(1); + var callback = manager.runUntilCompleteOrSuspend(new CompletableFuture()) + .handle((value, error) -> { + try { + CompletableFuture.runAsync(() -> manager.registerActiveThread("callback-lock-proof")) + .get(3, TimeUnit.SECONDS); + } catch (Exception failure) { + throw new AssertionError("Failure callbacks must not hold the activeThreads monitor", failure); + } + callbacks.countDown(); + return error; + }); + var owner = mock(BaseDurableOperation.class); + var ownerCompletion = new CompletableFuture(); + when(owner.getCompletionFuture()).thenReturn(ownerCompletion); + var unrelated = mock(BaseDurableOperation.class); + var unrelatedCompletion = new CompletableFuture(); + when(unrelated.getOperationId()).thenReturn("unrelated"); + when(unrelated.getCompletionFuture()).thenReturn(unrelatedCompletion); + manager.registerOperation(unrelated); + var buildingFailure = new CountDownLatch(1); + var releaseFailure = new CountDownLatch(1); + RuntimeException failure = mode.equals("reject") + ? new RejectedExecutionException("coordinator rejected") + : new IllegalArgumentException("continuation failed"); + var publisher = CompletableFuture.runAsync(() -> { + // A scheduling gate inside ordinary control construction, after the old caller's isClosing precheck. + // It does not replace manager state, futures, close, or stopAllOperations. + try (var model = mockStatic(ErrorObject.class, CALLS_REAL_METHODS)) { + model.when(ErrorObject::builder).thenAnswer(call -> { + buildingFailure.countDown(); + assertTrue(releaseFailure.await(3, TimeUnit.SECONDS)); + return call.callRealMethod(); + }); + if (mode.equals("reject")) { + assertSame( + failure, + assertThrows( + RejectedExecutionException.class, + () -> manager.runCheckpointContinuation( + owner, () -> fail("Rejected work"), task -> { + throw failure; + }))); + } else { + var observed = manager.runCheckpointContinuation( + owner, + () -> { + throw failure; + }, + Runnable::run); + assertSame( + failure, + assertThrows(CompletionException.class, observed::join) + .getCause()); + } + } + }); + CompletableFuture closing = null; + try { + assertTrue(buildingFailure.await(3, TimeUnit.SECONDS)); + if (closeFirst) { + closing = CompletableFuture.runAsync(manager::close); + assertInstanceOf( + SuspendExecutionException.class, + assertThrows(ExecutionException.class, () -> ownerCompletion.get(3, TimeUnit.SECONDS)) + .getCause()); + assertFalse(unrelatedCompletion.isDone()); + var pendingClose = closing; + assertThrows(TimeoutException.class, () -> pendingClose.get(100, TimeUnit.MILLISECONDS)); + } + releaseFailure.countDown(); + publisher.get(3, TimeUnit.SECONDS); + if (closing == null) closing = CompletableFuture.runAsync(manager::close); + closing.get(3, TimeUnit.SECONDS); + System.out.println("CONTINUATION_CLOSE_RACE mode=" + mode + " closeFirst=" + closeFirst + + " unrelatedStopped=" + unrelatedCompletion.isDone() + + " managerFailure=" + manager.isExecutionCompletedExceptionally()); + assertEquals( + !closeFirst, + unrelatedCompletion.isDone(), + "Only failure selected before close may stop unrelated operations"); + assertEquals(!closeFirst, manager.isExecutionCompletedExceptionally()); + if (!closeFirst) { + assertNotNull(callback.get(3, TimeUnit.SECONDS)); + assertEquals(0, callbacks.getCount()); + } else assertEquals(1, callbacks.getCount()); + assertDoesNotThrow(() -> manager.deregisterActiveThread("root")); + } finally { + releaseFailure.countDown(); + publisher.get(3, TimeUnit.SECONDS); + if (closing == null) closing = CompletableFuture.runAsync(manager::close); + closing.get(3, TimeUnit.SECONDS); + } + } +} diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/DurableExecutionTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/DurableExecutionTest.java index e0773aaeb..df72229ec 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/execution/DurableExecutionTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/DurableExecutionTest.java @@ -411,7 +411,12 @@ void testExecutorNotShutdownAfterMultipleHandlerInvocations() { (userInput, ctx) -> ctx.step("test1", String.class, stepCtx -> "Result 1: " + userInput), config); - assertEquals(ExecutionStatus.SUCCEEDED, output1.status()); + assertEquals( + ExecutionStatus.SUCCEEDED, + output1.status(), + () -> output1.error() == null + ? "No error payload" + : output1.error().errorType() + ": " + output1.error().errorMessage()); assertFalse(sharedExecutor.isShutdown(), "Executor should not be shutdown after first execution"); // Create second input with different execution operation @@ -440,7 +445,12 @@ void testExecutorNotShutdownAfterMultipleHandlerInvocations() { (userInput, ctx) -> ctx.step("test2", String.class, stepCtx -> "Result 2: " + userInput), config); - assertEquals(ExecutionStatus.SUCCEEDED, output2.status()); + assertEquals( + ExecutionStatus.SUCCEEDED, + output2.status(), + () -> output2.error() == null + ? "No error payload" + : output2.error().errorType() + ": " + output2.error().errorMessage()); assertFalse(sharedExecutor.isShutdown(), "Executor should not be shutdown after second execution"); // Verify both executions completed successfully and used the same executor 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 1552d2b21..872a69181 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 @@ -12,13 +12,23 @@ import java.util.List; import java.util.Map; import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.Executors; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.ValueSource; import software.amazon.awssdk.services.lambda.model.CheckpointUpdatedExecutionState; +import software.amazon.awssdk.services.lambda.model.ErrorObject; import software.amazon.awssdk.services.lambda.model.GetDurableExecutionStateResponse; import software.amazon.awssdk.services.lambda.model.Operation; import software.amazon.awssdk.services.lambda.model.OperationStatus; @@ -27,6 +37,7 @@ import software.amazon.lambda.durable.TestUtils; import software.amazon.lambda.durable.client.DurableExecutionClient; import software.amazon.lambda.durable.context.DurableContextImpl; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; import software.amazon.lambda.durable.model.DurableExecutionInput; import software.amazon.lambda.durable.model.OperationIdentifier; import software.amazon.lambda.durable.model.OperationSubType; @@ -50,6 +61,134 @@ private ExecutionManager createManager(List operations) { null); } + @SuppressWarnings("removal") + @ParameterizedTest + @ValueSource(strings = {"success", "ordinary", "vm", "death"}) + void bodySelectedBeforeManagerTerminationRetainsItsOutcome(String kind) { + try (var manager = createManager(List.of(executionOp()))) { + var body = new CompletableFuture(); + var selected = manager.runUntilCompleteOrSuspend(body); + Throwable failure = + switch (kind) { + case "ordinary" -> new IllegalArgumentException("body failure"); + case "vm" -> new InternalError("body fatal"); + case "death" -> new ThreadDeath(); + default -> null; + }; + if (failure == null) body.complete("body success"); + else body.completeExceptionally(failure); + var later = new UnrecoverableDurableExecutionException( + ErrorObject.builder().errorMessage("later manager failure").build(), true); + assertSame( + later, + assertThrows( + UnrecoverableDurableExecutionException.class, () -> manager.terminateExecution(later))); + if (failure == null) assertEquals("body success", selected.join()); + else + assertSame( + failure, + assertThrows(CompletionException.class, selected::join).getCause()); + } + } + + @SuppressWarnings("removal") + @ParameterizedTest + @CsvSource({"false,false", "false,true", "true,false", "true,true"}) + void continuationFatalSettlesOriginalBeforeEscapingWorker(boolean death, boolean rootActive) throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("coordinator fatal"); + var escaped = new CountDownLatch(1); + var release = new CountDownLatch(1); + var ownerFailure = new AtomicReference(); + var observation = new AtomicReference>(); + var settledAtEscape = new AtomicBoolean(); + var executor = Executors.newSingleThreadExecutor(task -> { + var thread = new Thread(task, "continuation-fatal-probe"); + thread.setUncaughtExceptionHandler((owner, failure) -> { + ownerFailure.set(failure); + settledAtEscape.set(observation.get().isDone()); + escaped.countDown(); + }); + return thread; + }); + try (var manager = createManager(List.of(executionOp(), stepOp("pending", OperationStatus.PENDING)))) { + if (rootActive) manager.registerActiveThread("root-probe"); + var future = manager.runCheckpointContinuation( + () -> { + try { + assertTrue(release.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + throw new AssertionError(interrupted); + } + throw fatal; + }, + executor); + observation.set(future); + var closedFromObserver = new AtomicBoolean(); + future.whenComplete((ignored, failure) -> { + manager.close(); // Must not wait for this callback's own continuation registration. + closedFromObserver.set(true); + }); + release.countDown(); + assertSame( + fatal, + assertThrows(ExecutionException.class, () -> future.get(3, TimeUnit.SECONDS)) + .getCause()); + assertTrue(escaped.await(2, TimeUnit.SECONDS), "Fatal must escape the actual coordinator worker"); + assertSame(fatal, ownerFailure.get()); + assertTrue(settledAtEscape.get(), "Observation must settle before worker escape"); + assertTrue(closedFromObserver.get(), "The activity lease must be released before observer callbacks"); + } finally { + release.countDown(); + executor.shutdownNow(); + assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @SuppressWarnings("removal") + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void inlineFatalContinuationReleasesItsRegistrationOnlyOnce(boolean death) { + Error fatal = death ? new ThreadDeath() : new InternalError("inline coordinator fatal"); + try (var manager = createManager(List.of(executionOp()))) { + manager.registerActiveThread("root"); + assertSame( + fatal, + assertThrows( + Error.class, + () -> manager.runCheckpointContinuation( + () -> { + throw fatal; + }, + Runnable::run))); + assertTrue(manager.isExecutionCompletedExceptionally()); + } + } + + @Test + void ordinaryContinuationFailureRetainsItsExistingObservationPolicy() throws Exception { + var original = new IllegalArgumentException("ordinary continuation"); + var worker = new AtomicReference(); + var executor = Executors.newSingleThreadExecutor(); + try (var manager = createManager(List.of(executionOp()))) { + manager.registerActiveThread("root"); + var result = manager.runCheckpointContinuation( + () -> { + worker.set(Thread.currentThread()); + throw original; + }, + executor); + assertSame( + original, + assertThrows(ExecutionException.class, () -> result.get(3, TimeUnit.SECONDS)) + .getCause()); + assertFalse(manager.isExecutionCompletedExceptionally()); + assertSame(worker.get(), executor.submit(Thread::currentThread).get(3, TimeUnit.SECONDS)); + } finally { + executor.shutdownNow(); + assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + private Operation executionOp() { return Operation.builder() .id(EXECUTION_OP_ID) @@ -372,6 +511,258 @@ void deferredSuspensionOccursWhenCheckpointFinishesWithoutReactivatingThread() { assertTrue(manager.isExecutionCompletedExceptionally()); } + @ParameterizedTest + @ValueSource(strings = {"complete", "register-worker", "fail"}) + void pollingContinuationRetainsActivityUntilItFinishesDispatch(String outcome) throws Exception { + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + var entered = new CountDownLatch(1); + var release = new CountDownLatch(1); + var caller = Thread.currentThread(); + var failure = new IllegalStateException("continuation failure"); + manager.registerActiveThread("root"); + var completion = manager.runCheckpointContinuation(() -> { + assertNotSame(caller, Thread.currentThread()); + entered.countDown(); + try { + assertTrue(release.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new AssertionError(interrupted); + } + if (outcome.equals("register-worker")) manager.registerActiveThread("next"); + if (outcome.equals("fail")) throw failure; + }); + try { + assertTrue(entered.await(3, TimeUnit.SECONDS)); + manager.deregisterActiveThread("root"); + assertFalse(manager.isExecutionCompletedExceptionally()); + assertFalse(completion.isDone()); + } finally { + release.countDown(); + } + if (outcome.equals("fail")) { + assertSame( + failure, + assertThrows(ExecutionException.class, () -> completion.get(3, TimeUnit.SECONDS)) + .getCause()); + } else completion.get(3, TimeUnit.SECONDS); + if (outcome.equals("register-worker")) { + assertFalse(manager.isExecutionCompletedExceptionally()); + assertThrows(SuspendExecutionException.class, () -> manager.deregisterActiveThread("next")); + } + assertTrue(manager.isExecutionCompletedExceptionally()); + manager.runCheckpointContinuation(() -> fail("No continuation may start after suspension")) + .get(3, TimeUnit.SECONDS); + } + + @Test + void cancellingQueuedContinuationDoesNotSkipItsWorkOrLeakActivity() { + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + var queued = new AtomicReference(); + var calls = new AtomicInteger(); + manager.registerActiveThread("root"); + var completion = manager.runCheckpointContinuation(calls::incrementAndGet, queued::set); + assertTrue(completion.cancel(false)); + manager.deregisterActiveThread("root"); + assertFalse(manager.isExecutionCompletedExceptionally(), "Queued work still owns the activity lease"); + queued.get().run(); + assertAll( + () -> assertEquals(1, calls.get()), + () -> assertTrue(manager.isExecutionCompletedExceptionally(), "The actual runnable releases its lease"), + () -> assertTrue(completion.isCancelled())); + } + + @Test + void cancellingRunningContinuationDoesNotReleaseActivityBeforeItsCleanup() throws Exception { + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + var queued = new AtomicReference(); + var entered = new CountDownLatch(1); + var release = new CountDownLatch(1); + manager.registerActiveThread("root"); + var completion = manager.runCheckpointContinuation( + () -> { + entered.countDown(); + try { + assertTrue(release.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new AssertionError(interrupted); + } + }, + queued::set); + var actualWorker = CompletableFuture.runAsync(queued.get()); + try { + assertTrue(entered.await(3, TimeUnit.SECONDS)); + assertTrue(completion.cancel(false)); + manager.deregisterActiveThread("root"); + assertFalse(manager.isExecutionCompletedExceptionally(), "Cancellation must not release running work"); + } finally { + release.countDown(); + } + actualWorker.get(3, TimeUnit.SECONDS); + assertTrue(manager.isExecutionCompletedExceptionally()); + assertTrue(completion.isCancelled()); + } + + @SuppressWarnings("removal") + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void coordinatorAdmissionFatalSelectsRetryBeforeLeaseRelease(boolean death) throws Exception { + Error fatal = death ? new ThreadDeath() : new InternalError("coordinator admission fatal"); + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + manager.registerActiveThread("root"); + var owner = mock(BaseDurableOperation.class); + when(owner.getCompletionFuture()).thenReturn(new CompletableFuture<>()); + var selected = manager.runUntilCompleteOrSuspend(new CompletableFuture<>()); + assertSame( + fatal, + assertThrows( + Error.class, + () -> manager.runCheckpointContinuation(owner, () -> fail("No accepted task"), task -> { + throw fatal; + }))); + var retry = assertInstanceOf( + UnrecoverableDurableExecutionException.class, + assertThrows(ExecutionException.class, () -> selected.get(3, TimeUnit.SECONDS)) + .getCause()); + assertTrue(retry.isRetryable()); + assertSame(fatal, retry.getCause()); + CompletableFuture.runAsync(manager::close).get(3, TimeUnit.SECONDS); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void ownedFailureReachesExecutionDespiteCancelledObservation(boolean cancelled) throws Exception { + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + manager.registerActiveThread("root"); + var owner = mock(BaseDurableOperation.class); + when(owner.getCompletionFuture()).thenReturn(new CompletableFuture<>()); + var queued = new AtomicReference(); + var failure = new IllegalArgumentException("owned continuation failure"); + var selected = manager.runUntilCompleteOrSuspend(new CompletableFuture<>()); + var observation = manager.runCheckpointContinuation( + owner, + () -> { + throw failure; + }, + queued::set); + if (cancelled) assertTrue(observation.cancel(false)); + queued.get().run(); + var control = assertInstanceOf( + UnrecoverableDurableExecutionException.class, + assertThrows(ExecutionException.class, () -> selected.get(3, TimeUnit.SECONDS)) + .getCause()); + assertTrue(control.isRetryable()); + assertSame(failure, control.getCause()); + if (cancelled) assertTrue(observation.isCancelled()); + else + assertSame( + failure, + assertThrows(ExecutionException.class, () -> observation.get(3, TimeUnit.SECONDS)) + .getCause()); + CompletableFuture.runAsync(manager::close).get(3, TimeUnit.SECONDS); + } + + @Test + void ownedCoordinatorRejectionSettlesExecutionAndReleasesAdmission() throws Exception { + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + manager.registerActiveThread("root"); + var owner = mock(BaseDurableOperation.class); + when(owner.getCompletionFuture()).thenReturn(new CompletableFuture<>()); + var selected = manager.runUntilCompleteOrSuspend(new CompletableFuture<>()); + var rejection = new RejectedExecutionException("coordinator admission rejected"); + assertSame( + rejection, + assertThrows( + RejectedExecutionException.class, + () -> manager.runCheckpointContinuation(owner, () -> fail("not admitted"), task -> { + throw rejection; + }))); + var control = assertInstanceOf( + UnrecoverableDurableExecutionException.class, + assertThrows(ExecutionException.class, () -> selected.get(3, TimeUnit.SECONDS)) + .getCause()); + assertTrue(control.isRetryable()); + assertSame(rejection, control.getCause()); + CompletableFuture.runAsync(manager::close).get(3, TimeUnit.SECONDS); + } + + @Test + void rejectedContinuationReleasesActivityAndPreservesTheRejection() { + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + var rejection = new RejectedExecutionException("coordinator unavailable"); + manager.registerActiveThread("root"); + assertSame( + rejection, + assertThrows( + RejectedExecutionException.class, + () -> manager.runCheckpointContinuation(() -> fail("Rejected work must not run"), task -> { + throw rejection; + }))); + assertThrows(SuspendExecutionException.class, () -> manager.deregisterActiveThread("root")); + } + + @ParameterizedTest + @ValueSource(strings = {"queued", "cancelled-observation", "rejected"}) + void admissionRacingCloseDrainsActualWorkAndPreservesUnrelatedOperations(String mode) throws Exception { + var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); + manager.registerActiveThread("root"); + var owner = mock(BaseDurableOperation.class); + var ownerCompletion = new CompletableFuture(); + when(owner.getCompletionFuture()).thenReturn(ownerCompletion); + var unrelated = mock(BaseDurableOperation.class); + var unrelatedCompletion = new CompletableFuture(); + when(unrelated.getOperationId()).thenReturn("unrelated"); + when(unrelated.getCompletionFuture()).thenReturn(unrelatedCompletion); + manager.registerOperation(unrelated); + var dispatchEntered = new CountDownLatch(1); + var releaseDispatch = new CountDownLatch(1); + var queued = new AtomicReference(); + var rejection = new RejectedExecutionException("rejected during close"); + var admission = CompletableFuture.supplyAsync(() -> + manager.runCheckpointContinuation(owner, () -> fail("Closing must stop this admitted body"), task -> { + dispatchEntered.countDown(); + try { + assertTrue(releaseDispatch.await(3, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new AssertionError(interrupted); + } + if (mode.equals("rejected")) throw rejection; + queued.set(task); + })); + assertTrue(dispatchEntered.await(3, TimeUnit.SECONDS)); + var closing = CompletableFuture.runAsync(manager::close); + try { + assertInstanceOf( + SuspendExecutionException.class, + assertThrows(ExecutionException.class, () -> ownerCompletion.get(3, TimeUnit.SECONDS)) + .getCause()); + assertThrows(TimeoutException.class, () -> closing.get(100, TimeUnit.MILLISECONDS)); + } finally { + releaseDispatch.countDown(); + } + if (mode.equals("rejected")) { + assertSame( + rejection, + assertThrows(ExecutionException.class, () -> admission.get(3, TimeUnit.SECONDS)) + .getCause()); + } else { + var observation = admission.get(3, TimeUnit.SECONDS); + if (mode.equals("cancelled-observation")) assertTrue(observation.cancel(false)); + assertThrows(TimeoutException.class, () -> closing.get(100, TimeUnit.MILLISECONDS)); + queued.get().run(); + } + closing.get(3, TimeUnit.SECONDS); + assertFalse(unrelatedCompletion.isDone(), "Normal close must not apply global stopAllOperations"); + assertFalse(manager.isExecutionCompletedExceptionally(), "Cleanup must not replace the selected root outcome"); + manager.runCheckpointContinuation( + () -> fail("Post-close admission must be rejected"), + task -> fail("Post-close work must not reach the executor")) + .get(3, TimeUnit.SECONDS); + assertDoesNotThrow(() -> manager.deregisterActiveThread("root")); + } + @Test void checkpointDeliveryIsAtomicWithOperationRegistration() throws Exception { var manager = createManager(List.of(executionOp(), stepOp("step", OperationStatus.PENDING))); diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/InvocationLifecycleTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/InvocationLifecycleTest.java new file mode 100644 index 000000000..cb9259472 --- /dev/null +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/InvocationLifecycleTest.java @@ -0,0 +1,414 @@ +// 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.ArgumentMatchers.any; +import static org.mockito.Mockito.*; +import static software.amazon.lambda.durable.TypeToken.get; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import org.slf4j.MDC; +import org.slf4j.helpers.BasicMDCAdapter; +import org.slf4j.spi.MDCAdapter; +import software.amazon.awssdk.services.lambda.model.CheckpointUpdatedExecutionState; +import software.amazon.awssdk.services.lambda.model.ErrorObject; +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.context.BaseContextImpl; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.model.ExecutionStatus; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; +import software.amazon.lambda.durable.plugin.InvocationInfo; +import software.amazon.lambda.durable.plugin.InvocationStatus; +import software.amazon.lambda.durable.serde.JacksonSerDes; +import software.amazon.lambda.durable.serde.SerDes; + +class InvocationLifecycleTest { + private MDCAdapter originalMdcAdapter; + + @BeforeEach + void installMdcAdapter() throws ReflectiveOperationException { + originalMdcAdapter = MDC.getMDCAdapter(); + setMdcAdapter(new BasicMDCAdapter()); + } + + @AfterEach + void restoreMdcAdapter() throws ReflectiveOperationException { + MDC.clear(); + BaseContextImpl.setCurrentContext(null); + setMdcAdapter(originalMdcAdapter); + } + + private static void setMdcAdapter(MDCAdapter adapter) throws ReflectiveOperationException { + var field = MDC.class.getDeclaredField("MDC_ADAPTER"); + field.setAccessible(true); + field.set(null, adapter); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void precompletedHandlerStillEndsOnItsOwner(boolean direct) throws Exception { + var workers = Executors.newSingleThreadExecutor(); + var plugin = new RecordingPlugin(); + var handlerThread = new AtomicReference(); + var executor = executor(task -> { + if (direct) task.run(); + else awaitTask(workers, task); + }); + try { + var output = DurableExecutor.execute( + input(), + null, + get(String.class), + (input, context) -> { + handlerThread.set(Thread.currentThread()); + return "done"; + }, + config(executor, plugin)); + assertEquals(ExecutionStatus.SUCCEEDED, output.status()); + assertEquals(1, plugin.starts.size()); + assertEquals(1, plugin.ends.size()); + assertSame(handlerThread.get(), plugin.startThreads.get(0)); + assertSame(handlerThread.get(), plugin.endThreads.get(0)); + assertEquals("invocation", plugin.endLocalValues.get(0)); + if (!direct) assertNotSame(Thread.currentThread(), handlerThread.get()); + } finally { + stop(workers); + } + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void reusedWorkerRestoresAmbientMdcAfterBothHooks(boolean fails) throws Exception { + var workers = Executors.newSingleThreadExecutor(); + var plugin = new RecordingPlugin(); + var ambient = Map.of("worker", "ambient"); + var callerMdc = MDC.getCopyOfContextMap(); + MDC.put("caller", "untouched"); + try { + workers.submit(() -> MDC.setContextMap(ambient)).get(5, TimeUnit.SECONDS); + var config = config(workers, plugin); + for (int invocation = 0; invocation < 2; invocation++) { + var output = DurableExecutor.execute( + input(), + null, + get(String.class), + (input, context) -> { + if (fails) throw new IllegalStateException("handler failure"); + return "done"; + }, + config); + assertEquals(fails ? ExecutionStatus.FAILED : ExecutionStatus.SUCCEEDED, output.status()); + assertEquals(ambient, workers.submit(MDC::getCopyOfContextMap).get(5, TimeUnit.SECONDS)); + assertNull(workers.submit(plugin.local::get).get(5, TimeUnit.SECONDS)); + assertEquals("untouched", MDC.get("caller")); + } + assertEquals(2, plugin.ends.size()); + assertEquals(List.of("invocation", "invocation"), plugin.endLocalValues); + assertEquals(plugin.startThreads, plugin.endThreads); + assertEquals(List.of(ambient, ambient), plugin.startMdc); + } finally { + if (callerMdc == null) MDC.clear(); + else MDC.setContextMap(callerMdc); + stop(workers); + } + } + + @Test + void inputDeserializationFailurePairsHooksOnWorkerWithoutRunningHandler() throws Exception { + var workers = Executors.newSingleThreadExecutor(); + var plugin = new RecordingPlugin(); + var handlerCalled = new AtomicBoolean(); + var serDes = mock(SerDes.class); + when(serDes.deserialize(any(), any())).thenThrow(new IllegalArgumentException("invalid input")); + var config = config(workers, plugin).toBuilder().withSerDes(serDes).build(); + try { + var output = DurableExecutor.execute( + input(), + null, + get(String.class), + (input, context) -> { + handlerCalled.set(true); + return "unreachable"; + }, + config); + assertEquals(ExecutionStatus.FAILED, output.status()); + assertFalse(handlerCalled.get()); + assertEquals(1, plugin.starts.size()); + assertNull(plugin.starts.get(0).executionInput()); + assertEquals(1, plugin.ends.size()); + assertEquals(InvocationStatus.FAILED, plugin.ends.get(0).invocationStatus()); + assertSame(plugin.startThreads.get(0), plugin.endThreads.get(0)); + assertNotSame(Thread.currentThread(), plugin.endThreads.get(0)); + assertEquals(List.of("invocation"), plugin.endLocalValues); + } finally { + stop(workers); + } + } + + @Test + void mdcCaptureFailureBeforeStartDoesNotDispatchEitherHook() throws Exception { + var workers = Executors.newSingleThreadExecutor(); + var plugin = new RecordingPlugin(); + var handlerCalled = new AtomicBoolean(); + var executor = executor(task -> workers.execute(() -> { + try (var mdc = mockStatic(MDC.class, CALLS_REAL_METHODS)) { + mdc.when(MDC::getCopyOfContextMap).thenThrow(new IllegalStateException("capture failed")); + task.run(); + } + })); + try { + var output = DurableExecutor.execute( + input(), + null, + get(String.class), + (input, context) -> { + handlerCalled.set(true); + return "unreachable"; + }, + config(executor, plugin)); + assertEquals(ExecutionStatus.FAILED, output.status()); + assertFalse(handlerCalled.get()); + assertTrue(plugin.starts.isEmpty()); + assertTrue(plugin.ends.isEmpty()); + } finally { + stop(workers); + } + } + + @Test + void callerCannotReturnUntilInvocationEndCompletes() throws Exception { + var enteredEnd = new CountDownLatch(1); + var releaseEnd = new CountDownLatch(1); + var plugin = new RecordingPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + enteredEnd.countDown(); + await(releaseEnd); + super.onInvocationEnd(info); + } + }; + var callers = Executors.newSingleThreadExecutor(); + var workers = Executors.newSingleThreadExecutor(); + try { + var response = callers.submit(() -> DurableExecutor.execute( + input(), null, get(String.class), (input, context) -> "done", config(workers, plugin))); + assertTrue(enteredEnd.await(5, TimeUnit.SECONDS)); + assertThrows(TimeoutException.class, () -> response.get(700, TimeUnit.MILLISECONDS)); + releaseEnd.countDown(); + assertEquals( + ExecutionStatus.SUCCEEDED, response.get(5, TimeUnit.SECONDS).status()); + assertEquals(1, plugin.ends.size()); + assertEquals(plugin.startThreads, plugin.endThreads); + } finally { + releaseEnd.countDown(); + stop(callers); + stop(workers); + } + } + + @Test + void callerCannotReturnUntilAmbientMdcIsRestored() throws Exception { + var restoring = new CountDownLatch(1); + var releaseRestore = new CountDownLatch(1); + var restored = new AtomicBoolean(); + var ambient = Map.of("worker", "ambient"); + var workers = Executors.newSingleThreadExecutor(); + var callers = Executors.newSingleThreadExecutor(); + var plugin = new RecordingPlugin(); + var executor = executor(task -> workers.execute(() -> { + MDC.setContextMap(ambient); + try (var mdc = mockStatic(MDC.class, CALLS_REAL_METHODS)) { + mdc.when(() -> MDC.setContextMap(ambient)).thenAnswer(call -> { + restoring.countDown(); + await(releaseRestore); + call.callRealMethod(); + restored.set(true); + return null; + }); + task.run(); + } + })); + try { + var response = callers.submit(() -> DurableExecutor.execute( + input(), null, get(String.class), (input, context) -> "done", config(executor, plugin))); + assertTrue(restoring.await(5, TimeUnit.SECONDS)); + assertThrows(TimeoutException.class, () -> response.get(700, TimeUnit.MILLISECONDS)); + releaseRestore.countDown(); + assertEquals( + ExecutionStatus.SUCCEEDED, response.get(5, TimeUnit.SECONDS).status()); + assertTrue(restored.get()); + assertEquals(1, plugin.ends.size()); + } finally { + releaseRestore.countDown(); + stop(callers); + stop(workers); + } + } + + @Test + void serializationFailureStillEndsOnTheHandlerThread() throws Exception { + var failure = new IllegalStateException("output cannot serialize"); + var serDes = spy(new JacksonSerDes()); + when(serDes.serialize("done")).thenThrow(failure); + var plugin = new RecordingPlugin(); + var workers = Executors.newSingleThreadExecutor(); + var config = config(workers, plugin).toBuilder().withSerDes(serDes).build(); + try { + assertSame( + failure, + assertThrows( + IllegalStateException.class, + () -> DurableExecutor.execute( + input(), null, get(String.class), (input, context) -> "done", config))); + assertEquals(1, plugin.starts.size()); + assertEquals(1, plugin.ends.size()); + assertEquals(plugin.startThreads, plugin.endThreads); + assertEquals(InvocationStatus.RETRYING, plugin.ends.get(0).invocationStatus()); + assertSame(failure, plugin.ends.get(0).executionError()); + assertNull(plugin.ends.get(0).executionResult()); + } finally { + stop(workers); + } + } + + @Test + void retryableLargeResultCheckpointFailureEndsWithOriginalRetryError() throws Exception { + var failure = new UnrecoverableDurableExecutionException( + ErrorObject.builder().errorMessage("retry delivery").build(), true); + var client = TestUtils.createMockClient(); + when(client.checkpoint(any(), any(), any())).thenThrow(failure); + var plugin = new RecordingPlugin(); + var workers = Executors.newSingleThreadExecutor(); + var config = config(workers, plugin).toBuilder() + .withDurableExecutionClient(client) + .build(); + try { + assertSame( + failure, + assertThrows( + UnrecoverableDurableExecutionException.class, + () -> DurableExecutor.execute( + input(), + null, + get(String.class), + (input, context) -> "x".repeat(7 * 1024 * 1024), + config))); + assertEquals(1, plugin.starts.size()); + assertEquals(1, plugin.ends.size()); + assertEquals(plugin.startThreads, plugin.endThreads); + assertEquals(InvocationStatus.RETRYING, plugin.ends.get(0).invocationStatus()); + assertSame(failure, plugin.ends.get(0).executionError()); + assertNull(plugin.ends.get(0).executionResult()); + } finally { + stop(workers); + } + } + + private static DurableConfig config(ExecutorService executor, DurableExecutionPlugin plugin) { + return DurableConfig.builder() + .withDurableExecutionClient(TestUtils.createMockClient()) + .withExecutorService(executor) + .withPlugins(plugin) + .build(); + } + + private static ExecutorService executor(Consumer dispatch) { + var executor = mock(ExecutorService.class); + doAnswer(call -> { + dispatch.accept(call.getArgument(0)); + return null; + }) + .when(executor) + .execute(any(Runnable.class)); + return executor; + } + + private static void awaitTask(ExecutorService workers, Runnable task) { + try { + workers.submit(task).get(5, TimeUnit.SECONDS); + } catch (Exception failure) { + throw new AssertionError(failure); + } + } + + private static void await(CountDownLatch latch) { + try { + if (!latch.await(5, TimeUnit.SECONDS)) throw new AssertionError("latch not released"); + } catch (InterruptedException failure) { + Thread.currentThread().interrupt(); + throw new AssertionError(failure); + } + } + + private static void stop(ExecutorService executor) throws InterruptedException { + executor.shutdownNow(); + assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS)); + } + + private static DurableExecutionInput input() { + var execution = Operation.builder() + .id("execution") + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.parse("2026-08-15T00:00:00Z")) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/test/execution", + "token", + CheckpointUpdatedExecutionState.builder().operations(execution).build()); + } + + private static class RecordingPlugin implements DurableExecutionPlugin { + final ThreadLocal local = new ThreadLocal<>(); + final List starts = new ArrayList<>(); + final List ends = new ArrayList<>(); + final List startThreads = new ArrayList<>(); + final List endThreads = new ArrayList<>(); + final List endLocalValues = new ArrayList<>(); + final List> startMdc = new ArrayList<>(); + + @Override + public void onInvocationStart(InvocationInfo info) { + starts.add(info); + startThreads.add(Thread.currentThread()); + startMdc.add(MDC.getCopyOfContextMap()); + local.set("invocation"); + MDC.put("plugin", "invocation"); + } + + @Override + public void onInvocationEnd(InvocationEndInfo info) { + ends.add(info); + endThreads.add(Thread.currentThread()); + endLocalValues.add(local.get()); + local.remove(); + MDC.put("plugin", "ended"); + } + } +} 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 6c923994f..39ea05e8d 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 @@ -284,6 +284,8 @@ void getFailedWithNullErrorDataThrowsStepFailedException() { @Test void replayPendingPollsAndResumesCheckLoop() throws Exception { + when(executionManager.runCheckpointContinuation(any(), any(Runnable.class))) + .thenAnswer(call -> CompletableFuture.runAsync(call.getArgument(1))); var pendingOp = Operation.builder() .id(OPERATION_ID) .name(OPERATION_NAME) diff --git a/sdk/src/test/java/software/amazon/lambda/durable/operation/WaitForConditionTerminalPollingTest.java b/sdk/src/test/java/software/amazon/lambda/durable/operation/WaitForConditionTerminalPollingTest.java new file mode 100644 index 000000000..ecac9a10b --- /dev/null +++ b/sdk/src/test/java/software/amazon/lambda/durable/operation/WaitForConditionTerminalPollingTest.java @@ -0,0 +1,119 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.operation; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.Mockito.*; + +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.*; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.junit.jupiter.params.provider.NullAndEmptySource; +import org.junit.jupiter.params.provider.ValueSource; +import software.amazon.awssdk.services.lambda.model.Operation; +import software.amazon.awssdk.services.lambda.model.OperationType; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.config.WaitForConditionConfig; +import software.amazon.lambda.durable.context.DurableContextImpl; +import software.amazon.lambda.durable.exception.IllegalDurableOperationException; +import software.amazon.lambda.durable.execution.ExecutionManager; +import software.amazon.lambda.durable.model.OperationIdentifier; +import software.amazon.lambda.durable.model.OperationSubType; +import software.amazon.lambda.durable.model.WaitForConditionResult; + +class WaitForConditionTerminalPollingTest { + @ParameterizedTest + @CsvSource({ + "CANCELLED,false", + "CANCELLED,true", + "TIMED_OUT,false", + "TIMED_OUT,true", + "STOPPED,false", + "STOPPED,true" + }) + void terminalObservationNeverRegistersAnotherPoll(String status, boolean alreadyKnown) throws Exception { + var fixture = new Fixture(alreadyKnown ? status : "PENDING"); + var result = fixture.poll(); + if (!alreadyKnown) fixture.deliver(status); + assertEquals(status, result.get(1, TimeUnit.SECONDS).statusAsString()); + assertEquals(alreadyKnown ? 0 : 1, fixture.polls.size()); + } + + @ParameterizedTest + @ValueSource(strings = {"NEW_UNKNOWN_STATUS"}) + @NullAndEmptySource + void malformedStatusTerminatesInsteadOfPollingForever(String status) throws Exception { + var fixture = new Fixture("PENDING"); + var result = fixture.poll(); + fixture.deliver(status); + var failure = assertThrows(ExecutionException.class, () -> result.get(1, TimeUnit.SECONDS)); + assertInstanceOf(IllegalDurableOperationException.class, failure.getCause()); + verify(fixture.manager).terminateExecution((IllegalDurableOperationException) failure.getCause()); + assertEquals(1, fixture.polls.size()); + } + + @Test + void nonterminalUpdatesStillWaitForReady() throws Exception { + var fixture = new Fixture("PENDING"); + var result = fixture.poll(); + fixture.deliver("STARTED"); + assertFalse(result.isDone()); + fixture.deliver("PENDING"); + assertFalse(result.isDone()); + fixture.deliver("READY"); + assertEquals("READY", result.get(1, TimeUnit.SECONDS).statusAsString()); + assertEquals(3, fixture.polls.size()); + } + + private static final class Fixture { + final ExecutionManager manager = mock(ExecutionManager.class); + final AtomicReference known; + final List> polls = new ArrayList<>(); + final WaitForConditionOperation operation; + + Fixture(String status) { + known = new AtomicReference<>(snapshot(status)); + var context = mock(DurableContextImpl.class); + when(context.getExecutionManager()).thenReturn(manager); + when(context.getDurableConfig()).thenReturn(DurableConfig.builder().build()); + when(manager.getOperationAndUpdateReplayState("condition")).thenAnswer(call -> known.get()); + when(manager.pollForOperationUpdates("condition")).thenAnswer(call -> { + var poll = new CompletableFuture(); + polls.add(poll); + return poll; + }); + operation = new WaitForConditionOperation<>( + OperationIdentifier.of("condition", "condition", OperationSubType.WAIT_FOR_CONDITION), + (state, contextIgnored) -> WaitForConditionResult.stopPolling(state), + TypeToken.get(Integer.class), + WaitForConditionConfig.builder().initialState(1).build(), + context); + } + + @SuppressWarnings("unchecked") + CompletableFuture poll() throws Exception { + var method = WaitForConditionOperation.class.getDeclaredMethod("pollUntilReady"); + method.setAccessible(true); + return (CompletableFuture) method.invoke(operation); + } + + void deliver(String status) { + var update = snapshot(status); + known.set(update); + polls.get(polls.size() - 1).complete(update); + } + + private static Operation snapshot(String status) { + return Operation.builder() + .id("condition") + .type(OperationType.STEP) + .status(status) + .build(); + } + } +} diff --git a/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginLinkageErrorTest.java b/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginLinkageErrorTest.java index 276e60738..5611215fb 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginLinkageErrorTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginLinkageErrorTest.java @@ -41,7 +41,10 @@ void otherErrorsRetainTheirExistingPropagation( plugin(healthyCalls::incrementAndGet))); assertSame(failure, assertThrows(Error.class, () -> dispatch.accept(runner))); - assertEquals(0, healthyCalls.get(), "Fatal and unrelated errors must not be blanket-caught"); + assertEquals( + hook.equals("invocation end") ? 1 : 0, + healthyCalls.get(), + "Invocation End must finish cleanup before propagating errors; other hooks stop immediately"); } static Stream linkageFailures() { diff --git a/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginRunnerTest.java b/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginRunnerTest.java index 82f10f927..c006b6252 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginRunnerTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/plugin/PluginRunnerTest.java @@ -155,7 +155,7 @@ void fireAndForget_callsAllHookTypes() { // ─── Awaited hooks ─────────────────────────────────────────────────── @Test - void awaitedHooks_callAllPlugins() { + void invocationEnd_unwindsPluginsInReverseRegistrationOrder() { var calls = new ArrayList(); var plugin1 = new TestPlugin("p1", calls); var plugin2 = new TestPlugin("p2", calls); @@ -163,7 +163,35 @@ void awaitedHooks_callAllPlugins() { runner.onInvocationEnd(invocationEndInfo()); - assertEquals(List.of("p1:onInvocationEnd", "p2:onInvocationEnd"), calls); + assertEquals(List.of("p2:onInvocationEnd", "p1:onInvocationEnd"), calls); + } + + @SuppressWarnings("removal") + @Test + void invocationEnd_runsRemainingCleanupThenRethrowsTheFirstError() { + for (var firstFailure : + List.of(new InternalError("fatal"), new AssertionError("assertion"), new ThreadDeath())) { + var calls = new ArrayList(); + var inner = new TestPlugin("inner", calls) { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + super.onInvocationEnd(info); + throw firstFailure; + } + }; + var laterFailure = new AssertionError("later cleanup failure"); + var middle = new TestPlugin("middle", calls) { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + super.onInvocationEnd(info); + throw laterFailure; + } + }; + var runner = new PluginRunner(List.of(new TestPlugin("outer", calls), middle, inner)); + assertSame(firstFailure, assertThrows(Error.class, () -> runner.onInvocationEnd(invocationEndInfo()))); + assertEquals(List.of(laterFailure), List.of(firstFailure.getSuppressed())); + assertEquals(List.of("inner:onInvocationEnd", "middle:onInvocationEnd", "outer:onInvocationEnd"), calls); + } } @Test @@ -171,7 +199,7 @@ void awaitedHooks_swallowExceptions_butCallRemainingPlugins() { var calls = new ArrayList(); var throwingPlugin = new ThrowingPlugin(); var normalPlugin = new TestPlugin("p2", calls); - var runner = new PluginRunner(List.of(throwingPlugin, normalPlugin)); + var runner = new PluginRunner(List.of(normalPlugin, throwingPlugin)); assertDoesNotThrow(() -> runner.onInvocationEnd(invocationEndInfo())); assertEquals(List.of("p2:onInvocationEnd"), calls);