diff --git a/.github/scripts/install_otel_test_wheels.py b/.github/scripts/install_otel_test_wheels.py new file mode 100644 index 000000000..b1332e5c2 --- /dev/null +++ b/.github/scripts/install_otel_test_wheels.py @@ -0,0 +1,104 @@ +"""Install and verify built artifacts without silently using editable sources.""" + +from __future__ import annotations + +import argparse +import ast +import hashlib +import importlib.metadata +import json +import subprocess +import sys +import zipfile +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[2] +CORE = "aws-durable-execution-sdk-python" +OTEL = CORE + "-otel" +TESTING = CORE + "-testing" + + +def built_wheel(package: str) -> Path: + directory = ROOT / "packages" / package + module = package.replace("-", "_") + about = ast.parse((directory / "src" / module / "__about__.py").read_text()) + version = next( + ast.literal_eval(node.value) + for node in about.body + if isinstance(node, ast.Assign) + and any( + isinstance(target, ast.Name) and target.id == "__version__" + for target in node.targets + ) + ) + wheels = list((directory / "dist").glob(f"{module}-{version}-*.whl")) + if len(wheels) != 1: + raise ValueError(f"Build exactly one {package} {version} wheel first: {wheels}") + return wheels[0] + + +def verify(wheel: Path, package: str) -> None: + installed = importlib.metadata.distribution(package) + direct = json.loads(installed.read_text("direct_url.json") or "{}") + assert not direct.get("dir_info", {}).get("editable"), direct + digest = hashlib.sha256(wheel.read_bytes()).hexdigest() + assert direct["archive_info"]["hashes"]["sha256"] == digest, direct + module = package.replace("-", "_") + with zipfile.ZipFile(wheel) as archive: + sources = [ + name + for name in archive.namelist() + if name.startswith(module + "/") and name.endswith(".py") + ] + assert sources + for name in sources: + path = Path(installed.locate_file(name)).resolve() + assert "site-packages" in path.parts, path + assert path.read_bytes() == archive.read(name), path + print( + json.dumps( + { + "package": package, + "version": installed.version, + "wheel": str(wheel), + "sha256": digest, + "verified_sources": len(sources), + } + ) + ) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--legacy-plugin", action="store_true") + args = parser.parse_args() + # This PR also repairs the local service simulator. Validate its actual + # built artifact with the SDK/plugin, retaining every test assertion. + packages = [CORE] if args.legacy_plugin else [CORE, OTEL, TESTING] + wheels = [built_wheel(package) for package in packages] + subprocess.run( + [ + sys.executable, + "-m", + "pip", + "install", + "--no-index", + "--no-deps", + "--force-reinstall", + *map(str, wheels), + ], + check=True, + ) + for wheel, package in zip(wheels, packages, strict=True): + verify(wheel, package) + if args.legacy_plugin: + assert importlib.metadata.version(OTEL) == "1.0.0" + import aws_durable_execution_sdk_python_otel as otel + + assert "site-packages" in Path(otel.__file__).resolve().parts + subprocess.run([sys.executable, "-m", "pip", "check"], check=True) + + +if __name__ == "__main__": + main() diff --git a/.github/scripts/tests/test_opentelemetry_conformance_workflow.py b/.github/scripts/tests/test_opentelemetry_conformance_workflow.py index 034be671e..24b55be8c 100644 --- a/.github/scripts/tests/test_opentelemetry_conformance_workflow.py +++ b/.github/scripts/tests/test_opentelemetry_conformance_workflow.py @@ -1,8 +1,12 @@ +import re from pathlib import Path import yaml +SHARED_WORKFLOW_REF = "bdb4f1cd0f9252c1aaa978bb8b341b71f2b9d9dc" +CONFORMANCE_TEST_REF = "75987d46a915bc37409eed3ea9c3617a924c9756" + WORKFLOW_PATH = ( Path(__file__).parents[2] / "workflows" / "opentelemetry-conformance-tests.yml" ) @@ -18,7 +22,10 @@ def test_opentelemetry_conformance_caller_uses_current_workflow_contract() -> No "uses: aws/aws-durable-execution-conformance-tests/.github/workflows/" "opentelemetry-orchestrator.yml@" ) - assert orchestrator in workflow + pinned_ref = re.search(re.escape(orchestrator) + r"([0-9a-f]{40})", workflow) + assert pinned_ref is not None + assert pinned_ref.group(1) == SHARED_WORKFLOW_REF + assert f"default: {CONFORMANCE_TEST_REF}" in workflow assert "python-opentelemetry.yml@" not in workflow assert "\n otlp_endpoint:" not in workflow @@ -28,7 +35,7 @@ def test_opentelemetry_conformance_caller_uses_current_workflow_contract() -> No "resource_prefix: p", "sdk_repository: aws/aws-durable-execution-sdk-python", "sdk_ref: ${{ github.event.pull_request.head.sha || github.sha }}", - "conformance_test_ref: ${{ inputs.conformance_test_ref || 'main' }}", + f"conformance_test_ref: ${{{{ inputs.conformance_test_ref || '{CONFORMANCE_TEST_REF}' }}}}", "checkout_sdk: true", f"examples_dir: {EXAMPLES_DIR}", "adot_release_repository: aws-observability/aws-otel-python-instrumentation", diff --git a/.github/tests/otel_lifecycle_compatibility_test.py b/.github/tests/otel_lifecycle_compatibility_test.py new file mode 100644 index 000000000..8e9ce857b --- /dev/null +++ b/.github/tests/otel_lifecycle_compatibility_test.py @@ -0,0 +1,162 @@ +"""Exercise real installed version pairs through the public durable handler.""" + +from __future__ import annotations + +import contextvars +from datetime import UTC, datetime +from importlib.metadata import version +import os +from pathlib import Path +import threading +from types import SimpleNamespace + +import pytest +from aws_durable_execution_sdk_python import durable_execution +from aws_durable_execution_sdk_python import execution as core_execution +from aws_durable_execution_sdk_python.lambda_service import ( + ExecutionDetails, + Operation, + OperationStatus, + OperationType, +) +from aws_durable_execution_sdk_python.plugin import DurableInstrumentationPlugin +from aws_durable_execution_sdk_python_otel.execution_plugin import ExecutionOtelPlugin +from aws_durable_execution_sdk_python_otel.invocation_plugin import InvocationOtelPlugin +from aws_durable_execution_sdk_python_otel.otel_plugin_config import OtelPluginConfig +from opentelemetry import baggage, context, trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from packaging.version import Version + + +class NoNetworkClient: + def __getattr__(self, name: str): + raise AssertionError(f"Unexpected service access: {name}") + + +@pytest.mark.parametrize("view", [InvocationOtelPlugin, ExecutionOtelPlugin]) +@pytest.mark.parametrize("order", ["alone", "baggage-first", "baggage-last"]) +@pytest.mark.parametrize("ambient_present", [False, True]) +def test_installed_pair_preserves_context_and_documents_legacy_fallback( + view, order, ambient_present +): + legacy = os.environ.get("OTEL_COMPAT_LEGACY") == "1" + assert Version(version("aws-durable-execution-sdk-python")) >= Version("2.1.0") + assert version("aws-durable-execution-sdk-python-otel") == ( + "1.0.0" if legacy else "1.1.0" + ) + assert "site-packages" in Path(core_execution.__file__).resolve().parts + import aws_durable_execution_sdk_python_otel as installed_otel + + assert "site-packages" in Path(installed_otel.__file__).resolve().parts + + def run() -> None: + provider = TracerProvider() + exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer("installed-lifecycle") + phases = [] + worker_threads = [] + caller_thread = threading.get_ident() + + class BaggagePlugin(DurableInstrumentationPlugin): + def on_invocation_start(self, info): + self.token = context.attach(baggage.set_baggage("customer", "present")) + phases.append("baggage-start") + worker_threads.append(threading.get_ident()) + + def on_invocation_end(self, info): + context.detach(self.token) + phases.append("baggage-end") + worker_threads.append(threading.get_ident()) + + plugin = view( + OtelPluginConfig( + tracer_provider=provider, + enrich_logger=False, + context_extractor=lambda _: None, + ) + ) + plugins = { + "alone": [plugin], + "baggage-first": [BaggagePlugin(), plugin], + "baggage-last": [plugin, BaggagePlugin()], + }[order] + body_parents = [] + + @durable_execution(plugins=plugins, boto3_client=NoNetworkClient()) + def handler(event, durable_context): + phases.append("body") + worker_threads.append(threading.get_ident()) + assert baggage.get_baggage("incoming") == "keep" + assert baggage.get_baggage("customer") == ( + None if order == "alone" else "present" + ) + parent = trace.get_current_span().get_span_context() + body_parents.append(parent) + with tracer.start_as_current_span("customer-span"): + pass + return "ok" + + operation = Operation( + operation_id="installed", + operation_type=OperationType.EXECUTION, + status=OperationStatus.STARTED, + start_timestamp=datetime(2026, 10, 8, tzinfo=UTC), + execution_details=ExecutionDetails(input_payload="{}"), + ) + event = { + "DurableExecutionArn": "test-arn/installed", + "CheckpointToken": "token", + "InitialExecutionState": { + "Operations": [operation.to_json_dict()], + "NextMarker": "", + }, + } + lambda_context = SimpleNamespace( + aws_request_id="installed", + client_context=None, + identity=None, + _epoch_deadline_time_in_ms=0, + invoked_function_arn="test-arn", + tenant_id=None, + ) + ambient = tracer.start_span("host") if ambient_present else trace.INVALID_SPAN + host = baggage.set_baggage( + "incoming", "keep", trace.set_span_in_context(ambient) + ) + token = context.attach(host) + try: + for _ in range(2): + assert handler(event, lambda_context)["Status"] == "SUCCEEDED" + assert context.get_current() is host + assert caller_thread not in worker_threads + assert phases == ( + ["body"] * 2 + if order == "alone" + else ["baggage-start", "body", "baggage-end"] * 2 + ) + spans = exporter.get_finished_spans() + if legacy and view is InvocationOtelPlugin: + # The released plugin does not attach an Invocation fallback. + assert body_parents == [ambient.get_span_context()] * 2 + else: + name = "Invocation" if view is InvocationOtelPlugin else "Workflow" + contexts = [span.context for span in spans if span.name == name] + assert contexts + assert all( + parent.is_valid and parent in contexts for parent in body_parents + ) + users = [span for span in spans if span.name == "customer-span"] + assert len(users) == 2 + assert [span.parent for span in users] == [ + parent if parent.is_valid else None for parent in body_parents + ] + assert plugin._context_tokens == {} + finally: + context.detach(token) + ambient.end() + provider.shutdown() + + contextvars.Context().run(run) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f3e682d08..f5efea050 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -10,7 +10,7 @@ on: branches: [ main ] pull_request: - branches: [ main ] + branches: ["main"] jobs: lint-commits: @@ -84,8 +84,6 @@ jobs: run: hatch run types:check - name: Run tests + coverage run: hatch run test:cov - - name: Verify supported legacy core compatibility - run: hatch run test-pypi-otel-legacy:test - name: Build distribution run: | for pkg in packages/*/; do @@ -96,6 +94,10 @@ jobs: cd "$GITHUB_WORKSPACE" fi done + - name: Test installed core and OTel wheels + run: hatch run test-wheel-otel:test + - name: Test released OTel with the new core wheel + run: hatch run test-wheel-otel-legacy:test - name: Verify OTel wheel dependency contract run: | OTEL_WHEEL=$(find packages/aws-durable-execution-sdk-python-otel/dist \ diff --git a/.github/workflows/cloud-tests.yml b/.github/workflows/cloud-tests.yml index 8e54be4fd..da80a43e2 100644 --- a/.github/workflows/cloud-tests.yml +++ b/.github/workflows/cloud-tests.yml @@ -133,7 +133,9 @@ jobs: echo "Could not resolve the latest ADOT Python layer for $AWS_REGION" exit 1 fi - aws lambda get-layer-version-by-arn \ + # Parallel jobs can throttle this read; keep retries local to the lookup. + AWS_RETRY_MODE=standard AWS_MAX_ATTEMPTS=8 \ + aws lambda get-layer-version-by-arn \ --arn "$ADOT_LAYER_ARN" \ --region "$AWS_REGION" \ --query LayerVersionArn \ diff --git a/.github/workflows/opentelemetry-conformance-tests.yml b/.github/workflows/opentelemetry-conformance-tests.yml index 216490be3..592b92603 100644 --- a/.github/workflows/opentelemetry-conformance-tests.yml +++ b/.github/workflows/opentelemetry-conformance-tests.yml @@ -45,7 +45,7 @@ on: conformance_test_ref: description: Conformance test commit SHA or branch name required: true - default: main + default: 75987d46a915bc37409eed3ea9c3617a924c9756 type: string # Backend stacks are shared across PRs. Queue whole runs so reusable @@ -66,14 +66,14 @@ 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@bdb4f1cd0f9252c1aaa978bb8b341b71f2b9d9dc with: language: python runs_on: codebuild-github-actions-runner-${{ github.run_id }}-${{ github.run_attempt }} resource_prefix: p sdk_repository: aws/aws-durable-execution-sdk-python sdk_ref: ${{ github.event.pull_request.head.sha || github.sha }} - conformance_test_ref: ${{ inputs.conformance_test_ref || 'main' }} + conformance_test_ref: ${{ inputs.conformance_test_ref || '75987d46a915bc37409eed3ea9c3617a924c9756' }} # Check the SDK out so the handlers and templates below are on disk. The handlers # themselves are installed from sdk_ref by src/requirements.txt during the SAM build. checkout_sdk: true diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 94176be2c..b583d7461 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -70,15 +70,25 @@ hatch run dev-otel:typecheck # type check otel only hatch run dev-examples:test # run examples tests only ``` -### PyPI release testing +### Installed package compatibility testing -To verify packages work against the published PyPI version of the core SDK (rather than the local workspace): +Build the core, OTel and testing-library distributions with `hatch build` in each package, then run +these commands from the repository root: ```bash -hatch run test-pypi-otel:test # test new OTel capabilities against capable installed core -hatch run test-pypi-otel-legacy:test # valid registrations/lifecycles on supported core 2.0.x -hatch run test-pypi-examples:test # test examples against PyPI core SDK -``` +hatch run test-wheel-otel:test # full OTel suite on the three built wheels +hatch run test-wheel-otel-legacy:test # released OTel 1.0.0 with the built core +hatch run test-pypi-examples:test # examples against the published core +``` + +The wheel lanes have no editable workspace members. They verify installed source +bytes and artifact hashes before exercising public handlers. OTel 1.1 requires +the redesigned core 2.1.0 lifecycle; it no longer claims compatibility with core +2.0.x. Publish core first. Building both wheels lets CI verify the intended pair +before that minimum is available on PyPI. The legacy-plugin lane documents the +actual core-only upgrade: host isolation is provided by the new core, while old +Invocation OTel does not gain the new fallback. Workspace tests continue to cover +the complete current implementation with `hatch run dev-otel:test`. ### Package-level commands diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/README.md b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/README.md index 5ff978c88..421de8f64 100644 --- a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/README.md +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/README.md @@ -33,7 +33,7 @@ template-long-running.yaml # otel-long-running suite tests/ # contract tests for the templates and handlers ``` -The 20 invocation and 20 execution requirements reuse the same scenario +The invocation and execution requirements reuse the same scenario handlers; the view is selected per function through the `OTEL_PLUGIN_MODE` environment variable, which `common.otel_plugin()` reads to pick `InvocationOtelPlugin` or `ExecutionOtelPlugin`. `template.yaml` deploys only the @@ -63,6 +63,12 @@ view named by its `OtelSuite` parameter. | `otel-invocation-18` | `otel_18_chained_invoke_failure.handler` | Verifies failed chained-invoke telemetry. | | `otel-invocation-19` | `otel_19_execution_failure.handler` | Verifies telemetry for a direct handler failure. | | `otel-invocation-20` | `otel_20_virtual_context.handler` | Verifies a virtual child-context span without context checkpoints. | +| `otel-invocation-21` | `otel_21_completed_step_replay.handler` | Replays a completed step after a successful wait/resume; each step body runs once. | +| `otel-invocation-22` | `otel_22_user_function_context.handler` | Creates ordinary user spans under the active handler, attempt, child, branch, and iteration contexts. | +| `otel-invocation-23` | `otel_23_callback_function_context.handler` | Verifies retry/check attempts, callback submitter, wrapped retry helper, and virtual-child callback parents. | +| `otel-invocation-24` | `otel_24_invocation_retry_status.handler` | Raises a retryable invocation error after a completed step, then resumes with its saved result. | +| `otel-invocation-25` | `otel_17_wait_for_callback_failure.handler` | Targets a failed callback without error details; service history must satisfy the explicit no-error-details precondition. | +| `otel-invocation-26` | `otel_26_external_callback_completion_replay.handler` | Completes a root callback after suspension, saves its result in a step, and replays it through two callback barriers. | | `otel-execution-1` | `otel_1_success.handler` | Verifies the execution-view workflow, step, and attempt hierarchy. | | `otel-execution-2` | `otel_2_wait_resume.handler` | Verifies the execution view across a resumed invocation. | | `otel-execution-3` | `otel_3_retry.handler` | Verifies the execution view across retry attempts. | @@ -83,6 +89,12 @@ view named by its `OtelSuite` parameter. | `otel-execution-18` | `otel_18_chained_invoke_failure.handler` | Verifies source and target failed workflow roots. | | `otel-execution-19` | `otel_19_execution_failure.handler` | Verifies a failed invocation without a completed workflow. | | `otel-execution-20` | `otel_20_virtual_context.handler` | Verifies a virtual child-context span under the workflow root. | +| `otel-execution-21` | `otel_21_completed_step_replay.handler` | Verifies completed-operation spans are exported once across normal successful replay. | +| `otel-execution-22` | `otel_22_user_function_context.handler` | Observes active execution-view callback contexts without supplying or repairing parents. | +| `otel-execution-23` | `otel_23_callback_function_context.handler` | Verifies the same SDK-owned callback lifecycle parents in the execution view. | +| `otel-execution-24` | `otel_24_invocation_retry_status.handler` | Verifies invocation retry status independently of step retry and recovery re-exports. | +| `otel-execution-25` | `otel_17_wait_for_callback_failure.handler` | Targets the same errorless failed callback and its `UNSET` leaf in the execution view. | +| `otel-execution-26` | `otel_26_external_callback_completion_replay.handler` | Requires one terminal root-callback export at first completion and no duplicate exports on two later replays. | | `otel-long-running-1` | `otel_long_running_1_wait.handler` | Verifies wait and resume telemetry across a long durable suspension. | | `otel-long-running-2` | `otel_long_running_2_retry.handler` | Verifies retry telemetry across a long durable backoff. | | `otel-long-running-3` | `otel_long_running_3_callback.handler` | Verifies callback telemetry when completion arrives after a long delay. | @@ -91,6 +103,55 @@ view named by its `OtelSuite` parameter. The runner discovers each mapping from `TestingMetadata.TestDescription` on the functions in the templates. +## External completion and status coverage + +The local testing library includes callback success, failure and timeout in the +next invocation's `UpdatedOperationIds`. The core delivers each terminal update +notification once per invocation, including when a resumed operation and a later +checkpoint response carry the same completion. It tracks actual notifications, +preserves first delivery when the update-ID metadata is absent, and clears that +tracking at invocation boundaries. Public runner regressions cover memory and +file stores, stored step results, failure payloads and two subsequent replays. +An exactly empty serialized callback error is represented as absent in both +stores, matching the SDK's service parser. Present fields remain intact, +including empty messages, types, data and stack lists; the enclosing failed +callback future still raises the same error as the file-store baseline. +Detailed callback-failure history retains the service's empty error `Payload` +object with `Truncated: false`, while SDK-facing state still has no error details. +Metadata-only history and nonempty errors retain their existing representation. + +Case 24 covers invocation `RETRY` becoming `RETRYING`/`UNSET`. Case 25 targets the +separate rule that a failed operation without error details remains `UNSET`. +Its cloud coverage requires both views to pass the raw service-history no-error-details +precondition and the telemetry assertions. Local file-store results establish +SDK behavior; that AWS service precondition still needs independent validation. +`CANCELLED`, `TIMED_OUT` and `STOPPED` without error details have explicit unit +coverage in both views; their corresponding cloud paths are not established. +The existing success/`OK` and detailed-error/`ERROR` controls remain in the suites. + +Case 26 revisits a root-level public callback on every replay. The driver waits +for `InvocationCompleted` before completing the target and each barrier. The +target's terminal span must precede the observed step and must not be exported +again during the two later resumes. Invocation view also retains its initial +pending callback segment. Cloud validation must count the raw S3 export records +without deduplicating equal span IDs; local runner checks alone do not establish +cloud coverage. + +## Callback coverage boundary + +Cases 22 and 23 create normal user spans from the active context, without an +explicit parent or a copied execution ARN. They cover handler context, step and +condition attempts, child and branch bodies, callback submission, wrapped +`with_retry` body/strategy callbacks, and a virtual child. Case 23 places the +retry helper and virtual child after callback completion so ordinary successful +replay does not repeat their probes. + +This does not promise an operation/attempt parent for every arbitrary callback. +General retry/wait policies, serializers, summary generators, and item naming +have phase-specific caller scopes outside the wrapped user-function lifecycle; +these cases do not impose a new ownership policy on them. Instrumentation +extensions and user-created threads are outside this business-callback contract. + ## How a handler maps to a requirement ```yaml diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/Makefile b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/Makefile index 320577e7f..8382a49cb 100644 --- a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/Makefile +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/Makefile @@ -5,6 +5,8 @@ build-Otel14MapFailure build-Otel15WaitInterrupted build-Otel16WaitForConditionFailure \ build-Otel17WaitForCallbackFailure build-Otel18ChainedInvokeFailure \ build-Otel18InvokeTarget build-Otel19ExecutionFailure build-Otel20VirtualContext \ + build-Otel21CompletedStepReplay build-Otel22UserFunctionContext build-Otel23CallbackFunctionContext build-Otel24InvocationRetryStatus \ + build-Otel25CallbackFailureWithoutError build-Otel26ExternalCallbackReplay \ build-OtelExecution1Success \ build-OtelExecution2WaitResume build-OtelExecution3Retry \ build-OtelExecution4TerminalFailure build-OtelExecution5ChildContext \ @@ -17,6 +19,9 @@ build-OtelExecution17WaitForCallbackFailure build-OtelExecution18ChainedInvokeFailure \ build-OtelExecution18InvokeTarget build-OtelExecution19ExecutionFailure \ build-OtelExecution20VirtualContext \ + build-OtelExecution21CompletedStepReplay build-OtelExecution22UserFunctionContext build-OtelExecution23CallbackFunctionContext \ + build-OtelExecution24InvocationRetryStatus \ + build-OtelExecution25CallbackFailureWithoutError build-OtelExecution26ExternalCallbackReplay \ build-OtelLongRunning1Wait \ build-OtelLongRunning2Retry build-OtelLongRunning3Callback \ build-OtelLongRunning4ChainedInvoke build-OtelLongRunning4InvokeTarget @@ -28,6 +33,8 @@ build-Otel11InvokeTarget build-Otel12ChildContextFailure build-Otel13ParallelFai build-Otel14MapFailure build-Otel15WaitInterrupted build-Otel16WaitForConditionFailure \ build-Otel17WaitForCallbackFailure build-Otel18ChainedInvokeFailure \ build-Otel18InvokeTarget build-Otel19ExecutionFailure build-Otel20VirtualContext \ +build-Otel21CompletedStepReplay build-Otel22UserFunctionContext build-Otel23CallbackFunctionContext build-Otel24InvocationRetryStatus \ +build-Otel25CallbackFailureWithoutError build-Otel26ExternalCallbackReplay \ build-OtelExecution1Success \ build-OtelExecution2WaitResume build-OtelExecution3Retry \ build-OtelExecution4TerminalFailure build-OtelExecution5ChildContext \ @@ -40,6 +47,9 @@ build-OtelExecution16WaitForConditionFailure \ build-OtelExecution17WaitForCallbackFailure build-OtelExecution18ChainedInvokeFailure \ build-OtelExecution18InvokeTarget build-OtelExecution19ExecutionFailure \ build-OtelExecution20VirtualContext \ +build-OtelExecution21CompletedStepReplay build-OtelExecution22UserFunctionContext build-OtelExecution23CallbackFunctionContext \ +build-OtelExecution24InvocationRetryStatus \ +build-OtelExecution25CallbackFailureWithoutError build-OtelExecution26ExternalCallbackReplay \ build-OtelLongRunning1Wait \ build-OtelLongRunning2Retry build-OtelLongRunning3Callback \ build-OtelLongRunning4ChainedInvoke build-OtelLongRunning4InvokeTarget: diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_21_completed_step_replay.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_21_completed_step_replay.py new file mode 100644 index 000000000..2ea9ef4a0 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_21_completed_step_replay.py @@ -0,0 +1,36 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Normal successful replay of a completed step for OTel case 21.""" + +from __future__ import annotations + +from typing import Any + +from aws_durable_execution_sdk_python import ( + DurableContext, + StepContext, + durable_execution, + durable_step, +) +from aws_durable_execution_sdk_python.config import Duration +from common import otel_plugin, require_scenario + + +@durable_step +def before_wait(_step_context: StepContext) -> str: + return "before" + + +@durable_step +def after_wait(_step_context: StepContext) -> str: + return "after" + + +@durable_execution(plugins=[otel_plugin()]) +def handler(event: dict[str, Any], context: DurableContext) -> str: + require_scenario(event, "completed-step-replay") + before = context.step(before_wait(), name="otel-before-wait") + context.wait(Duration.from_seconds(1), name="otel-replay-wait") + after = context.step(after_wait(), name="otel-after-wait") + return f"{before}-{after}" diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_22_user_function_context.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_22_user_function_context.py new file mode 100644 index 000000000..4f103ba4e --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_22_user_function_context.py @@ -0,0 +1,113 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Observe the real active SDK context inside public user callbacks.""" + +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any + +from aws_durable_execution_sdk_python import ( + DurableContext, + StepContext, + durable_execution, +) +from aws_durable_execution_sdk_python.config import ( + Duration, + MapConfig, + ParallelBranch, + ParallelConfig, +) +from common import otel_plugin, require_scenario +from opentelemetry import trace + + +def probe(label: str) -> None: + if not trace.get_current_span().get_span_context().is_valid: + raise RuntimeError(f"No active span for conformance.{label}") + span = trace.get_tracer("aws-durable-execution-conformance").start_span( + f"conformance.{label}", attributes={"conformance.callback": label} + ) + span.end() + + +def step_body(_context: StepContext) -> str: + probe("step") + return "step" + + +def child_step(_context: StepContext) -> str: + probe("child-step") + return "child-step" + + +def child_body(context: DurableContext) -> str: + probe("child") + context.step(child_step, name="otel-context-child-step") + probe("child-restored") + return "child" + + +def parallel_step_a(_context: StepContext) -> str: + probe("parallel-step-a") + return "a" + + +def parallel_step_b(_context: StepContext) -> str: + probe("parallel-step-b") + return "b" + + +def parallel_a(context: DurableContext) -> str: + probe("parallel-a") + return context.step(parallel_step_a, name="otel-context-branch-step-a") + + +def parallel_b(context: DurableContext) -> str: + probe("parallel-b") + return context.step(parallel_step_b, name="otel-context-branch-step-b") + + +def iteration_name(_item: int, index: int) -> str: + return ("otel-context-iteration-0", "otel-context-iteration-1")[index] + + +def mapper( + context: DurableContext, item: int, index: int, _items: Sequence[int] +) -> int: + probe(("map-0", "map-1")[index]) + + def map_step(_step_context: StepContext) -> int: + probe(("map-step-0", "map-step-1")[index]) + return item + + return context.step( + map_step, name=("otel-context-map-step-0", "otel-context-map-step-1")[index] + ) + + +@durable_execution(plugins=[otel_plugin()]) +def handler(event: dict[str, Any], context: DurableContext) -> str: + require_scenario(event, "user-function-context") + probe("handler") + context.step(step_body, name="otel-context-step") + context.run_in_child_context(child_body, name="otel-context-child") + context.parallel( + [ + ParallelBranch(parallel_a, name="otel-context-branch-a"), + ParallelBranch(parallel_b, name="otel-context-branch-b"), + ], + name="otel-context-parallel", + config=ParallelConfig(max_concurrency=2), + ) + context.map( + [0, 1], + mapper, + name="otel-context-map", + config=MapConfig(max_concurrency=2, item_namer=iteration_name), + ) + probe("handler-restored") + context.wait(Duration.from_seconds(1), name="otel-context-resume") + probe("handler-after-resume") + return "context-complete" diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_23_callback_function_context.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_23_callback_function_context.py new file mode 100644 index 000000000..c7fae30ec --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_23_callback_function_context.py @@ -0,0 +1,136 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Observe callback contexts at SDK-owned user-function lifecycle boundaries.""" + +from __future__ import annotations + +from typing import Any + +from aws_durable_execution_sdk_python import ( + DurableContext, + StepContext, + durable_execution, +) +from aws_durable_execution_sdk_python.config import ( + ChildConfig, + Duration, + JitterStrategy, + StepConfig, +) +from aws_durable_execution_sdk_python.exceptions import ChildContextError +from aws_durable_execution_sdk_python.retries import ( + RetryDecision, + RetryStrategyConfig, + WithRetryConfig, + create_retry_strategy, + with_retry, +) +from aws_durable_execution_sdk_python.types import ( + WaitForCallbackContext, + WaitForConditionCheckContext, +) +from aws_durable_execution_sdk_python.waits import ( + WaitForConditionConfig, + WaitForConditionDecision, +) +from common import otel_plugin, require_scenario +from opentelemetry import trace + + +HELPER_FAILURE = "intentional-helper-failure" + + +def probe(label: str) -> None: + if not trace.get_current_span().get_span_context().is_valid: + raise RuntimeError(f"No active span for conformance.{label}") + span = trace.get_tracer("aws-durable-execution-conformance").start_span( + f"conformance.{label}", attributes={"conformance.callback": label} + ) + span.end() + + +def retry_step(context: StepContext) -> str: + probe(f"retry-attempt-{context.attempt}") + if context.attempt == 1: + raise RuntimeError("intentional-step-retry") + return "retried" + + +def check_condition(state: int, _context: WaitForConditionCheckContext) -> int: + next_state = state + 1 + probe(f"condition-check-{next_state}") + return next_state + + +def wait_strategy(state: int, _attempt: int) -> WaitForConditionDecision: + if state >= 2: + return WaitForConditionDecision.stop_polling() + return WaitForConditionDecision.continue_waiting(Duration.from_seconds(1)) + + +def submit_callback(_callback_id: str, _context: WaitForCallbackContext) -> None: + probe("callback-submitter") + + +def helper_body(_context: DurableContext, _attempt: int) -> str: + probe("with-retry-body") + raise RuntimeError(HELPER_FAILURE) + + +def helper_retry_strategy(_error: Exception, _attempt: int) -> RetryDecision: + probe("with-retry-strategy") + return RetryDecision.no_retry() + + +def virtual_child(_context: DurableContext) -> str: + probe("virtual-child") + return "virtual" + + +@durable_execution(plugins=[otel_plugin()]) +def handler(event: dict[str, Any], context: DurableContext) -> str: + require_scenario(event, "callback-function-context") + context.step( + retry_step, + name="otel-context-retry-step", + config=StepConfig( + retry_strategy=create_retry_strategy( + RetryStrategyConfig( + max_attempts=2, + initial_delay=Duration.from_seconds(1), + max_delay=Duration.from_seconds(1), + backoff_rate=1.0, + jitter_strategy=JitterStrategy.NONE, + retryable_error_types=[RuntimeError], + ) + ) + ), + ) + context.wait_for_condition( + check_condition, + name="otel-context-condition", + config=WaitForConditionConfig(initial_state=0, wait_strategy=wait_strategy), + ) + context.wait_for_callback(submit_callback, name="otel-context-callback") + + # These callbacks run after the last asynchronous wait, so normal replay + # does not repeat the helper body or the checkpointless virtual child. + try: + with_retry( + context, + helper_body, + WithRetryConfig(retry_strategy=helper_retry_strategy), + name="otel-context-with-retry", + ) + except ChildContextError as error: + if error.message != HELPER_FAILURE: + raise + else: + raise RuntimeError("Expected the intentional helper failure") + context.run_in_child_context( + virtual_child, + name="otel-context-virtual", + config=ChildConfig(is_virtual=True), + ) + return "callback-context-complete" diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_24_invocation_retry_status.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_24_invocation_retry_status.py new file mode 100644 index 000000000..568622517 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_24_invocation_retry_status.py @@ -0,0 +1,33 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""A real invocation retry for OTel status-mapping case 24.""" + +from __future__ import annotations + +from typing import Any + +from aws_durable_execution_sdk_python import ( + DurableContext, + StepContext, + durable_execution, + durable_step, +) +from aws_durable_execution_sdk_python.exceptions import InvocationError +from common import otel_plugin, require_scenario + + +@durable_step +def before_invocation_retry(_step_context: StepContext) -> str: + return "saved" + + +@durable_execution(plugins=[otel_plugin()]) +def handler(event: dict[str, Any], context: DurableContext) -> str: + require_scenario(event, "invocation-retry-status") + # Capture this at entry: consuming the completed step can end replay mode. + entered_replay = context.is_replaying() + context.step(before_invocation_retry(), name="otel-before-invocation-retry") + if not entered_replay: + raise InvocationError("Conformance invocation retry after a completed step") + return "retry-complete" diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_26_external_callback_completion_replay.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_26_external_callback_completion_replay.py new file mode 100644 index 000000000..4e562e2f5 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/src/otel_26_external_callback_completion_replay.py @@ -0,0 +1,53 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Public callback completion followed by two controlled replay barriers.""" + +from __future__ import annotations + +from typing import Any, cast + +from aws_durable_execution_sdk_python import ( + DurableContext, + StepContext, + durable_execution, + durable_step, +) +from aws_durable_execution_sdk_python.config import ( + CallbackConfig, + WaitForCallbackConfig, +) +from aws_durable_execution_sdk_python.serdes import JsonSerDes +from aws_durable_execution_sdk_python.types import WaitForCallbackContext +from common import otel_plugin, require_scenario + + +def submit_callback(_callback_id: str, _context: WaitForCallbackContext) -> None: + return None + + +@durable_step +def observe_target(_context: StepContext, result: str) -> str: + return result + + +@durable_execution(plugins=[otel_plugin()]) +def handler(event: dict[str, Any], context: DurableContext) -> str: + require_scenario(event, "external-callback-completion-replay") + config = WaitForCallbackConfig(serdes=JsonSerDes()) + target = cast( + str, + context.create_callback( + name="otel-external-target", config=CallbackConfig(serdes=JsonSerDes()) + ).result(), + ) + observed = context.step( + observe_target(target), name="otel-external-target-observed" + ) + one = context.wait_for_callback( + submit_callback, name="otel-external-barrier-one", config=config + ) + two = context.wait_for_callback( + submit_callback, name="otel-external-barrier-two", config=config + ) + return "/".join((observed, one, two)) diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/template.yaml b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/template.yaml index be70e122d..8d898ba46 100644 --- a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/template.yaml +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/template.yaml @@ -409,6 +409,90 @@ Resources: FunctionName: !Sub "${AWS::StackName}-otel-invocation-20" Role: !Ref LambdaExecutionRoleArn + Otel21CompletedStepReplay: + Type: AWS::Serverless::Function + Condition: DeployInvocationView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-invocation-21 + Properties: + CodeUri: src/ + Handler: otel_21_completed_step_replay.handler + FunctionName: !Sub "${AWS::StackName}-otel-invocation-21" + Role: !Ref LambdaExecutionRoleArn + + Otel22UserFunctionContext: + Type: AWS::Serverless::Function + Condition: DeployInvocationView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-invocation-22 + Properties: + CodeUri: src/ + Handler: otel_22_user_function_context.handler + FunctionName: !Sub "${AWS::StackName}-otel-invocation-22" + Role: !Ref LambdaExecutionRoleArn + + Otel23CallbackFunctionContext: + Type: AWS::Serverless::Function + Condition: DeployInvocationView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-invocation-23 + Properties: + CodeUri: src/ + Handler: otel_23_callback_function_context.handler + FunctionName: !Sub "${AWS::StackName}-otel-invocation-23" + Role: !Ref LambdaExecutionRoleArn + + Otel24InvocationRetryStatus: + Type: AWS::Serverless::Function + Condition: DeployInvocationView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-invocation-24 + Properties: + CodeUri: src/ + Handler: otel_24_invocation_retry_status.handler + FunctionName: !Sub "${AWS::StackName}-otel-invocation-24" + Role: !Ref LambdaExecutionRoleArn + + Otel25CallbackFailureWithoutError: + Type: AWS::Serverless::Function + Condition: DeployInvocationView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-invocation-25 + Properties: + CodeUri: src/ + Handler: otel_17_wait_for_callback_failure.handler + FunctionName: !Sub "${AWS::StackName}-otel-invocation-25" + Role: !Ref LambdaExecutionRoleArn + + Otel26ExternalCallbackReplay: + Type: AWS::Serverless::Function + Condition: DeployInvocationView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-invocation-26 + Properties: + CodeUri: src/ + Handler: otel_26_external_callback_completion_replay.handler + FunctionName: !Sub "${AWS::StackName}-otel-invocation-26" + Role: !Ref LambdaExecutionRoleArn + OtelExecution1Success: Type: AWS::Serverless::Function Condition: DeployExecutionView @@ -781,3 +865,105 @@ Resources: Environment: Variables: OTEL_PLUGIN_MODE: execution + + OtelExecution21CompletedStepReplay: + Type: AWS::Serverless::Function + Condition: DeployExecutionView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-execution-21 + Properties: + CodeUri: src/ + Handler: otel_21_completed_step_replay.handler + FunctionName: !Sub "${AWS::StackName}-otel-execution-21" + Role: !Ref LambdaExecutionRoleArn + Environment: + Variables: + OTEL_PLUGIN_MODE: execution + + OtelExecution22UserFunctionContext: + Type: AWS::Serverless::Function + Condition: DeployExecutionView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-execution-22 + Properties: + CodeUri: src/ + Handler: otel_22_user_function_context.handler + FunctionName: !Sub "${AWS::StackName}-otel-execution-22" + Role: !Ref LambdaExecutionRoleArn + Environment: + Variables: + OTEL_PLUGIN_MODE: execution + + OtelExecution23CallbackFunctionContext: + Type: AWS::Serverless::Function + Condition: DeployExecutionView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-execution-23 + Properties: + CodeUri: src/ + Handler: otel_23_callback_function_context.handler + FunctionName: !Sub "${AWS::StackName}-otel-execution-23" + Role: !Ref LambdaExecutionRoleArn + Environment: + Variables: + OTEL_PLUGIN_MODE: execution + + OtelExecution24InvocationRetryStatus: + Type: AWS::Serverless::Function + Condition: DeployExecutionView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-execution-24 + Properties: + CodeUri: src/ + Handler: otel_24_invocation_retry_status.handler + FunctionName: !Sub "${AWS::StackName}-otel-execution-24" + Role: !Ref LambdaExecutionRoleArn + Environment: + Variables: + OTEL_PLUGIN_MODE: execution + + OtelExecution25CallbackFailureWithoutError: + Type: AWS::Serverless::Function + Condition: DeployExecutionView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-execution-25 + Properties: + CodeUri: src/ + Handler: otel_17_wait_for_callback_failure.handler + FunctionName: !Sub "${AWS::StackName}-otel-execution-25" + Role: !Ref LambdaExecutionRoleArn + Environment: + Variables: + OTEL_PLUGIN_MODE: execution + + OtelExecution26ExternalCallbackReplay: + Type: AWS::Serverless::Function + Condition: DeployExecutionView + Metadata: + BuildMethod: makefile + TestingMetadata: + TestDescription: + - otel-execution-26 + Properties: + CodeUri: src/ + Handler: otel_26_external_callback_completion_replay.handler + FunctionName: !Sub "${AWS::StackName}-otel-execution-26" + Role: !Ref LambdaExecutionRoleArn + Environment: + Variables: + OTEL_PLUGIN_MODE: execution diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_callback_failure_without_error.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_callback_failure_without_error.py new file mode 100644 index 000000000..a04a92c0d --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_callback_failure_without_error.py @@ -0,0 +1,270 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Public callback failures preserve error-detail semantics across stores.""" + +from __future__ import annotations + +import importlib.util +import json +import sys +import time +from pathlib import Path + +import pytest +from aws_durable_execution_sdk_python.lambda_service import ErrorObject +from aws_durable_execution_sdk_python.plugin import OperationEndInfo, OperationType +from aws_durable_execution_sdk_python_testing.model import ( + GetDurableExecutionHistoryResponse, + events_to_operations, +) +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from aws_durable_execution_sdk_python_testing.stores.filesystem import ( + FileSystemExecutionStore, +) +from opentelemetry import context +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, +) + + +SRC_DIR = Path(__file__).resolve().parents[1] / "src" + + +def _run_public_callback_failure( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], + filesystem: bool, + error_kind: str, +) -> dict[str, object]: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + callback_errors: list[ErrorObject | None] = [] + original_end = plugin.on_operation_end + + def observe_error(info: OperationEndInfo) -> None: + if info.operation_type is OperationType.CALLBACK: + callback_errors.append(info.error) + original_end(info) + + monkeypatch.setattr(plugin, "on_operation_end", observe_error) + original_context = context.get_current() + common_spec = importlib.util.spec_from_file_location( + "common", SRC_DIR / "common.py" + ) + assert common_spec is not None and common_spec.loader is not None + common = importlib.util.module_from_spec(common_spec) + monkeypatch.setitem(sys.modules, "common", common) + common_spec.loader.exec_module(common) + monkeypatch.setattr(common, "otel_plugin", lambda: plugin) + spec = importlib.util.spec_from_file_location( + "otel_17_wait_for_callback_failure", + SRC_DIR / "otel_17_wait_for_callback_failure.py", + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + error = { + "omitted": None, + "empty": ErrorObject(message=None, type=None, data=None, stack_trace=None), + "rich": ErrorObject( + message="explicit failure", + type="CallbackFailure", + data=None, + stack_trace=None, + ), + "empty-message": ErrorObject( + message="", type=None, data=None, stack_trace=None + ), + "empty-stack": ErrorObject(message=None, type=None, data=None, stack_trace=[]), + "empty-type": ErrorObject(message=None, type="", data=None, stack_trace=None), + "empty-data": ErrorObject(message=None, type=None, data="", stack_trace=None), + }[error_kind] + try: + store = FileSystemExecutionStore.create(tmp_path) if filesystem else None + with DurableFunctionTestRunner( + handler=module.handler, + store=store, + poll_interval=0.01, + execution_timeout=20, + skip_time=False, + ) as runner: + arn = runner.run_async( + input=json.dumps({"scenario": "wait-for-callback-failure"}) + ) + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + events = runner.get_execution_history( + arn, include_execution_data=True + ).events + starts = [ + event for event in events if event.event_type == "CallbackStarted" + ] + if starts and any( + event.event_type == "InvocationCompleted" + and event.event_id > starts[0].event_id + for event in events + ): + details = starts[0].callback_started_details + assert details is not None and details.callback_id is not None + callback_id = details.callback_id + break + time.sleep(0.01) + else: + raise AssertionError("Callback did not suspend before failure") + if error_kind == "omitted": + runner.send_callback_failure(callback_id) + else: + runner.send_callback_failure(callback_id, error=error) + result = runner.wait_for_result(arn, timeout=10) + with_data = runner.get_execution_history(arn, include_execution_data=True) + without_data = runner.get_execution_history( + arn, include_execution_data=False + ) + + assert result.status.value == "FAILED" + spans = exporter.get_finished_spans() + leaves = [ + span + for span in spans + if (span.attributes or {}).get("durable.operation.type") == "CALLBACK" + and (span.attributes or {}).get("durable.operation.status") == "FAILED" + ] + assert len(leaves) == 1 + has_details = error_kind not in {"omitted", "empty"} + assert [ + item.to_dict() if item is not None else None for item in callback_errors + ] == [error.to_dict() if has_details and error is not None else None] + if has_details and not filesystem: + assert callback_errors[0] is error + assert leaves[0].status.status_code.name == ( + "ERROR" if has_details else "UNSET" + ) + assert [event.name for event in leaves[0].events] == ( + ["exception"] if has_details else [] + ) + # The public failed future still raises, so its enclosing context fails. + parents = [ + span + for span in spans + if span.name == "otel-failed-callback" + and (span.attributes or {}).get("durable.operation.status") == "FAILED" + ] + assert len(parents) == 1 + assert parents[0].status.status_code.name == "ERROR" + invocations = sorted( + [span for span in spans if span.name == "Invocation"], + key=lambda span: span.start_time or 0, + ) + statuses = [ + (span.attributes or {}).get("durable.invocation.status") + for span in invocations + ] + assert statuses == ["PENDING", "FAILED"] + assert result.error is not None + assert context.get_current() == original_context + for history in (with_data, without_data): + wire = history.to_dict() + decoded = GetDurableExecutionHistoryResponse.from_dict(wire) + original_error = next( + event["CallbackFailedDetails"]["Error"] + for event in wire["Events"] + if event["EventType"] == "CallbackFailed" + ) + for candidate in (history, decoded): + assert ( + next( + event.to_dict()["CallbackFailedDetails"]["Error"] + for event in candidate.events + if event.event_type == "CallbackFailed" + ) + == original_error + ) + callback = next( + operation + for operation in events_to_operations(candidate.events) + if operation.callback_details is not None + ) + details = callback.callback_details + assert details is not None + assert ( + details.error.to_dict() if details.error is not None else None + ) == (error.to_dict() if has_details and error is not None else None) + return { + "status": result.status.value, + "caller_error": dict(result.error.to_dict()), + "invocation_statuses": statuses, + "submitted_error": dict(error.to_dict()) if error is not None else {}, + "history_error": next( + event.to_dict()["CallbackFailedDetails"]["Error"] + for event in with_data.events + if event.event_type == "CallbackFailed" + ), + "metadata_error": next( + event.to_dict()["CallbackFailedDetails"]["Error"] + for event in without_data.events + if event.event_type == "CallbackFailed" + ), + } + finally: + provider.shutdown() + + +@pytest.mark.parametrize("plugin_type", [ExecutionOtelPlugin, InvocationOtelPlugin]) +@pytest.mark.parametrize( + "error_kind", + [ + "omitted", + "empty", + "rich", + "empty-message", + "empty-stack", + "empty-type", + "empty-data", + ], +) +def test_public_callback_failure_preserves_error_details( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], + error_kind: str, +) -> None: + memory = _run_public_callback_failure( + monkeypatch, tmp_path / "memory", plugin_type, False, error_kind + ) + filesystem = _run_public_callback_failure( + monkeypatch, tmp_path / "filesystem", plugin_type, True, error_kind + ) + assert memory == filesystem + expected_payload = memory["submitted_error"] + for outcome in [memory, filesystem]: + # Actual AWS history retains an empty Payload object for this failure, + # independently of the SDK-facing absent callback error. + # Retain the existing flags for nonempty errors; only the observed + # no-details projection is changed here. + assert outcome["history_error"] == { + "Payload": expected_payload, + "Truncated": bool(expected_payload), + } + # Preserve this API's existing metadata-only projection. + assert outcome["metadata_error"] == ( + {"Payload": expected_payload, "Truncated": True} + if expected_payload + else {"Truncated": True} + ) diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_callback_function_context.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_callback_function_context.py new file mode 100644 index 000000000..2aecd8cc6 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_callback_function_context.py @@ -0,0 +1,122 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Verify actual SDK callback parents across retries and external completion.""" + +from __future__ import annotations + +import importlib.util +import json +import sys +import time +from collections import Counter +from pathlib import Path + +import pytest +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from opentelemetry import context, trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, +) + + +SRC_DIR = Path(__file__).resolve().parents[1] / "src" +EXPECTED_PARENTS = { + "retry-attempt-1": "otel-context-retry-step attempt 1", + "retry-attempt-2": "otel-context-retry-step attempt 2", + "condition-check-1": "otel-context-condition attempt 1", + "condition-check-2": "otel-context-condition attempt 2", + "callback-submitter": "otel-context-callback submitter attempt 1", + "with-retry-body": "otel-context-with-retry", + "with-retry-strategy": "otel-context-with-retry", + "virtual-child": "otel-context-virtual", +} + + +@pytest.mark.parametrize("plugin_type", [ExecutionOtelPlugin, InvocationOtelPlugin]) +def test_callback_probes_keep_their_actual_sdk_parent( + monkeypatch: pytest.MonkeyPatch, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + monkeypatch.setattr(trace, "get_tracer_provider", lambda: provider) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + before_context = context.get_current() + common_spec = importlib.util.spec_from_file_location( + "common", SRC_DIR / "common.py" + ) + assert common_spec is not None and common_spec.loader is not None + common = importlib.util.module_from_spec(common_spec) + monkeypatch.setitem(sys.modules, "common", common) + common_spec.loader.exec_module(common) + monkeypatch.setattr(common, "otel_plugin", lambda: plugin) + spec = importlib.util.spec_from_file_location( + "otel_23_callback_function_context", + SRC_DIR / "otel_23_callback_function_context.py", + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + try: + with DurableFunctionTestRunner(handler=module.handler) as runner: + arn = runner.run_async( + input=json.dumps({"scenario": "callback-function-context"}), timeout=30 + ) + callback_id = runner.wait_for_callback( + arn, name="otel-context-callback create callback id", timeout=10 + ) + # External delivery delay, outside the durable handler. + time.sleep(1) + runner.send_callback_success(callback_id, result=b"callback-complete") + result = runner.wait_for_result(arn, timeout=15) + assert result.status.value == "SUCCEEDED" + assert result.result is not None + assert json.loads(result.result) == "callback-context-complete" + spans = exporter.get_finished_spans() + probes = [span for span in spans if span.name.startswith("conformance.")] + assert Counter(span.name for span in probes) == Counter( + f"conformance.{label}" for label in EXPECTED_PARENTS + ) + spans_by_id = { + (span.context.trace_id, span.context.span_id): span + for span in spans + if span.context is not None + } + sdk_traces = { + span.context.trace_id + for span in spans + if span.context is not None + and span.attributes is not None + and span.attributes.get("durable.execution.arn") == arn + } + for probe_span in probes: + assert probe_span.context is not None + assert probe_span.parent is not None + assert probe_span.attributes is not None + label = str(probe_span.attributes["conformance.callback"]) + assert "durable.execution.arn" not in probe_span.attributes + assert probe_span.context.trace_id in sdk_traces + parent = spans_by_id[ + (probe_span.context.trace_id, probe_span.parent.span_id) + ] + assert parent.name == EXPECTED_PARENTS[label] + assert parent.attributes is not None + assert parent.attributes["durable.execution.arn"] == arn + assert context.get_current() == before_context + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_completed_step_replay.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_completed_step_replay.py new file mode 100644 index 000000000..d331decca --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_completed_step_replay.py @@ -0,0 +1,100 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Exercise case 21 through the public decorator and local durable runner.""" + +from __future__ import annotations + +import importlib.util +import json +import sys +from pathlib import Path +from typing import Any + +import pytest +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from opentelemetry import context +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, +) + + +SRC_DIR = Path(__file__).resolve().parents[1] / "src" + + +@pytest.mark.parametrize("plugin_type", [ExecutionOtelPlugin, InvocationOtelPlugin]) +def test_completed_step_replay_uses_saved_result( + monkeypatch: pytest.MonkeyPatch, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + before_context = context.get_current() + common_spec = importlib.util.spec_from_file_location( + "common", SRC_DIR / "common.py" + ) + assert common_spec is not None and common_spec.loader is not None + common = importlib.util.module_from_spec(common_spec) + monkeypatch.setitem(sys.modules, "common", common) + common_spec.loader.exec_module(common) + monkeypatch.setattr(common, "otel_plugin", lambda: plugin) + spec = importlib.util.spec_from_file_location( + "otel_21_completed_step_replay", SRC_DIR / "otel_21_completed_step_replay.py" + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + calls: list[str] = [] + + def track(factory: Any, label: str) -> Any: + def make_step() -> Any: + step = factory() + + def counted(step_context: Any) -> str: + calls.append(label) + return step(step_context) + + return counted + + return make_step + + monkeypatch.setattr(module, "before_wait", track(module.before_wait, "before")) + monkeypatch.setattr(module, "after_wait", track(module.after_wait, "after")) + try: + with DurableFunctionTestRunner(handler=module.handler) as runner: + result = runner.run( + input=json.dumps({"scenario": "completed-step-replay"}), timeout=15 + ) + assert result.status.value == "SUCCEEDED" + assert result.result is not None + assert json.loads(result.result) == "before-after" + assert calls == ["before", "after"] + spans = exporter.get_finished_spans() + invocations = [span for span in spans if span.name == "Invocation"] + assert [ + span.attributes["durable.invocation.status"] + for span in invocations + if span.attributes is not None + ] == ["PENDING", "SUCCEEDED"] + if plugin_type is ExecutionOtelPlugin: + # Count raw exports, including duplicates that reuse a span ID. + for name in ("otel-before-wait", "otel-replay-wait", "otel-after-wait"): + assert len([span for span in spans if span.name == name]) == 1 + assert context.get_current() == before_context + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_external_callback_completion_replay.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_external_callback_completion_replay.py new file mode 100644 index 000000000..2ec2f34fa --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_external_callback_completion_replay.py @@ -0,0 +1,195 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""A real callback completion is exported before two later public replays.""" + +from __future__ import annotations + +import importlib.util +import json +import sys +import time +from pathlib import Path +from typing import Any + +import pytest +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from aws_durable_execution_sdk_python_testing.stores.filesystem import ( + FileSystemExecutionStore, +) +from opentelemetry import context +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, +) + + +SRC_DIR = Path(__file__).resolve().parents[1] / "src" + + +def _suspended_callback( + runner: DurableFunctionTestRunner, arn: str, name: str +) -> tuple[str, tuple[int, int]]: + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + events = runner.get_execution_history(arn, include_execution_data=True).events + starts = [ + event + for event in events + if event.event_type == "CallbackStarted" and event.name == name + ] + completions = [ + event + for event in events + if event.event_type == "InvocationCompleted" + and starts + and event.event_id > starts[0].event_id + ] + if starts and completions: + details = starts[0].callback_started_details + assert details is not None and details.callback_id is not None + return details.callback_id, (starts[0].event_id, completions[0].event_id) + time.sleep(0.01) + raise AssertionError(f"Callback {name} did not suspend") + + +@pytest.mark.parametrize("plugin_type", [ExecutionOtelPlugin, InvocationOtelPlugin]) +def test_external_callback_completion_precedes_later_replays( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + before_context = context.get_current() + common_spec = importlib.util.spec_from_file_location( + "common", SRC_DIR / "common.py" + ) + assert common_spec is not None and common_spec.loader is not None + common = importlib.util.module_from_spec(common_spec) + monkeypatch.setitem(sys.modules, "common", common) + common_spec.loader.exec_module(common) + monkeypatch.setattr(common, "otel_plugin", lambda: plugin) + spec = importlib.util.spec_from_file_location( + "otel_26_external_callback_completion_replay", + SRC_DIR / "otel_26_external_callback_completion_replay.py", + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + calls: list[str] = [] + original_observe = module.observe_target + + def counted_observe(value: str) -> Any: + step = original_observe(value) + + def counted(step_context: Any) -> str: + calls.append(value) + return step(step_context) + + return counted + + monkeypatch.setattr(module, "observe_target", counted_observe) + gates: list[tuple[int, int]] = [] + try: + with DurableFunctionTestRunner( + handler=module.handler, + store=FileSystemExecutionStore.create(tmp_path), + poll_interval=0.01, + execution_timeout=25, + skip_time=False, + ) as runner: + arn = runner.run_async( + input=json.dumps({"scenario": "external-callback-completion-replay"}) + ) + for name, payload in ( + ("otel-external-target", "target"), + ("otel-external-barrier-one create callback id", "one"), + ("otel-external-barrier-two create callback id", "two"), + ): + callback_id, gate = _suspended_callback(runner, arn, name) + gates.append(gate) + runner.send_callback_success( + callback_id, result=json.dumps(payload).encode("utf-8") + ) + result = runner.wait_for_result(arn, timeout=10) + history = runner.get_execution_history(arn, include_execution_data=True) + + assert result.status.value == "SUCCEEDED" + assert result.result is not None + assert json.loads(result.result) == "target/one/two" + assert calls == ["target"] + assert gates == [(2, 3), (8, 11), (15, 18)] + assert [ + event.event_id + for event in history.events + if event.event_type == "InvocationCompleted" + ] == [3, 11, 18, 21] + assert history.events[-1].event_type == "ExecutionSucceeded" + assert history.events[-1].event_id == 22 + + spans = exporter.get_finished_spans() + assert all(span.context is not None for span in spans) + assert len({span.context.trace_id for span in spans if span.context}) == 1 + invocations = sorted( + [span for span in spans if span.name == "Invocation"], + key=lambda span: span.start_time or 0, + ) + assert len(invocations) == 4 + assert [ + (span.attributes or {}).get("durable.invocation.status") + for span in invocations + ] == ["PENDING", "PENDING", "PENDING", "SUCCEEDED"] + targets = [ + span + for span in spans + if span.name == "otel-external-target" + and (span.attributes or {}).get("durable.operation.status") == "SUCCEEDED" + ] + # Count raw exports, including duplicate records with identical span IDs. + assert len(targets) == 1 + target = targets[0] + assert target.status.status_code.name == "OK" + observed = [ + span + for span in spans + if span.name == "otel-external-target-observed attempt 1" + ] + assert len(observed) == 1 + assert target.end_time is not None and observed[0].start_time is not None + assert target.end_time <= observed[0].start_time + assert invocations[2].start_time is not None + assert invocations[3].start_time is not None + assert target.end_time < invocations[2].start_time < invocations[3].start_time + parent = ( + invocations[1] + if plugin_type is InvocationOtelPlugin + else next(span for span in spans if span.name == "Workflow") + ) + assert target.parent is not None and parent.context is not None + assert target.parent.span_id == parent.context.span_id + # Invocation view also exports the legitimate first-invocation segment. + pending_target = [ + span + for span in spans + if span.name == "otel-external-target" + and (span.attributes or {}).get("durable.operation.status") == "STARTED" + ] + assert len(pending_target) == (1 if plugin_type is InvocationOtelPlugin else 0) + assert context.get_current() == before_context + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_invocation_retry_status.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_invocation_retry_status.py new file mode 100644 index 000000000..d2851a709 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_invocation_retry_status.py @@ -0,0 +1,110 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Exercise a real invocation retry through the public decorator and local runner.""" + +from __future__ import annotations + +import importlib.util +import json +import sys +from pathlib import Path +from typing import Any + +import pytest +from aws_durable_execution_sdk_python.plugin import InvocationEndInfo, InvocationStatus +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from opentelemetry import context +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, +) + + +SRC_DIR = Path(__file__).resolve().parents[1] / "src" + + +@pytest.mark.parametrize("plugin_type", [ExecutionOtelPlugin, InvocationOtelPlugin]) +def test_invocation_retry_preserves_completed_step( + monkeypatch: pytest.MonkeyPatch, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + statuses: list[InvocationStatus] = [] + original_end = plugin.on_invocation_end + + def observe_end(info: InvocationEndInfo) -> None: + statuses.append(info.status) + original_end(info) + + monkeypatch.setattr(plugin, "on_invocation_end", observe_end) + before_context = context.get_current() + common_spec = importlib.util.spec_from_file_location( + "common", SRC_DIR / "common.py" + ) + assert common_spec is not None and common_spec.loader is not None + common = importlib.util.module_from_spec(common_spec) + monkeypatch.setitem(sys.modules, "common", common) + common_spec.loader.exec_module(common) + monkeypatch.setattr(common, "otel_plugin", lambda: plugin) + spec = importlib.util.spec_from_file_location( + "otel_24_invocation_retry_status", + SRC_DIR / "otel_24_invocation_retry_status.py", + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + calls: list[str] = [] + + def track(factory: Any, label: str) -> Any: + def make_step() -> Any: + step = factory() + + def counted(step_context: Any) -> str: + calls.append(label) + return step(step_context) + + return counted + + return make_step + + monkeypatch.setattr( + module, + "before_invocation_retry", + track(module.before_invocation_retry, "saved"), + ) + try: + with DurableFunctionTestRunner(handler=module.handler) as runner: + result = runner.run( + input=json.dumps({"scenario": "invocation-retry-status"}), timeout=15 + ) + assert result.status.value == "SUCCEEDED" + assert result.result is not None + assert json.loads(result.result) == "retry-complete" + assert calls == ["saved"] + assert statuses == [InvocationStatus.RETRY, InvocationStatus.SUCCEEDED] + spans = exporter.get_finished_spans() + invocations = [span for span in spans if span.name == "Invocation"] + assert len(invocations) == 2 + assert [span.status.status_code.name for span in invocations] == ["UNSET", "OK"] + workflows = [span for span in spans if span.name == "Workflow"] + assert len(workflows) == 1 + assert workflows[0].status.status_code.name == "OK" + assert context.get_current() == before_context + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_otel_examples.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_otel_examples.py index 5908dd7d9..61c0d24e1 100644 --- a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_otel_examples.py +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_otel_examples.py @@ -50,6 +50,12 @@ ("Otel18ChainedInvokeFailure", "otel-invocation-18"), ("Otel19ExecutionFailure", "otel-invocation-19"), ("Otel20VirtualContext", "otel-invocation-20"), + ("Otel21CompletedStepReplay", "otel-invocation-21"), + ("Otel22UserFunctionContext", "otel-invocation-22"), + ("Otel23CallbackFunctionContext", "otel-invocation-23"), + ("Otel24InvocationRetryStatus", "otel-invocation-24"), + ("Otel25CallbackFailureWithoutError", "otel-invocation-25"), + ("Otel26ExternalCallbackReplay", "otel-invocation-26"), ("OtelExecution1Success", "otel-execution-1"), ("OtelExecution2WaitResume", "otel-execution-2"), ("OtelExecution3Retry", "otel-execution-3"), @@ -70,6 +76,12 @@ ("OtelExecution18ChainedInvokeFailure", "otel-execution-18"), ("OtelExecution19ExecutionFailure", "otel-execution-19"), ("OtelExecution20VirtualContext", "otel-execution-20"), + ("OtelExecution21CompletedStepReplay", "otel-execution-21"), + ("OtelExecution22UserFunctionContext", "otel-execution-22"), + ("OtelExecution23CallbackFunctionContext", "otel-execution-23"), + ("OtelExecution24InvocationRetryStatus", "otel-execution-24"), + ("OtelExecution25CallbackFailureWithoutError", "otel-execution-25"), + ("OtelExecution26ExternalCallbackReplay", "otel-execution-26"), ] EXPECTED_LONG_RUNNING_MAPPINGS: list[tuple[str, str]] = [ ("OtelLongRunning1Wait", "otel-long-running-1"), @@ -152,6 +164,11 @@ "otel_18_chained_invoke_failure", "otel_19_execution_failure", "otel_20_virtual_context", + "otel_21_completed_step_replay", + "otel_22_user_function_context", + "otel_23_callback_function_context", + "otel_24_invocation_retry_status", + "otel_26_external_callback_completion_replay", "otel_long_running_1_wait", "otel_long_running_2_retry", "otel_long_running_3_callback", diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_user_function_context.py b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_user_function_context.py new file mode 100644 index 000000000..2b3cc26d6 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests-otel/tests/test_user_function_context.py @@ -0,0 +1,258 @@ +# SPDX-FileCopyrightText: 2026-present Amazon.com, Inc. or its affiliates. +# +# SPDX-License-Identifier: Apache-2.0 +"""Exercise case 22's callback contexts through real concurrent work and resume.""" + +from __future__ import annotations + +import importlib.util +import json +import sys +from collections import Counter +from contextlib import nullcontext +from pathlib import Path +from threading import Barrier +from typing import Any + +import pytest +from aws_durable_execution_sdk_python.execution import DurableExecutionInvocationInput +from aws_durable_execution_sdk_python.types import LambdaContext +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from opentelemetry import context, trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, +) + + +SRC_DIR = Path(__file__).resolve().parents[1] / "src" +CALLBACK_PARENTS = { + "step": "otel-context-step attempt 1", + "child": "otel-context-child", + "child-step": "otel-context-child-step attempt 1", + "child-restored": "otel-context-child", + "parallel-a": "otel-context-branch-a", + "parallel-step-a": "otel-context-branch-step-a attempt 1", + "parallel-b": "otel-context-branch-b", + "parallel-step-b": "otel-context-branch-step-b attempt 1", + "map-0": "otel-context-iteration-0", + "map-step-0": "otel-context-map-step-0 attempt 1", + "map-1": "otel-context-iteration-1", + "map-step-1": "otel-context-map-step-1 attempt 1", +} +HANDLER_COUNTS = {"handler": 2, "handler-restored": 2, "handler-after-resume": 1} + + +@pytest.mark.parametrize("plugin_type", [ExecutionOtelPlugin, InvocationOtelPlugin]) +@pytest.mark.parametrize( + "ambient", [False, True], ids=["no-ambient", "unrelated-ambient"] +) +def test_user_function_probes_keep_sdk_context_across_resume( + monkeypatch: pytest.MonkeyPatch, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], + ambient: bool, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + monkeypatch.setattr(trace, "get_tracer_provider", lambda: provider) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + before_context = context.get_current() + common_spec = importlib.util.spec_from_file_location( + "common", SRC_DIR / "common.py" + ) + assert common_spec is not None and common_spec.loader is not None + common = importlib.util.module_from_spec(common_spec) + monkeypatch.setitem(sys.modules, "common", common) + common_spec.loader.exec_module(common) + monkeypatch.setattr(common, "otel_plugin", lambda: plugin) + spec = importlib.util.spec_from_file_location( + "otel_22_user_function_context", SRC_DIR / "otel_22_user_function_context.py" + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + + barriers: dict[str, Barrier] = {} + for pair in ( + ("parallel-a", "parallel-b"), + ("parallel-step-a", "parallel-step-b"), + ("map-0", "map-1"), + ("map-step-0", "map-step-1"), + ): + barrier = Barrier(2, timeout=5) + barriers.update(dict.fromkeys(pair, barrier)) + original_probe = module.probe + + def synchronized_probe(label: str) -> None: + callback_context = context.get_current() + # Both sibling callbacks must be active before either observes its + # parent. Synchronize only in the test; never supply a tracing context. + if label in barriers: + barriers[label].wait() + original_probe(label) + if label in barriers: + barriers[label].wait() + assert context.get_current() is callback_context + + monkeypatch.setattr(module, "probe", synchronized_probe) + invocation_contexts: list[tuple[context.Context, context.Context]] = [] + + def entry( + event: DurableExecutionInvocationInput, lambda_context: LambdaContext + ) -> dict[str, Any]: + # Simulate host instrumentation outside the unmodified durable handler. + # With no extracted parent, its unrelated trace must not become the + # parent of any SDK operation or handler probe. + host_context = context.get_current() + scope = ( + provider.get_tracer("host-instrumentation").start_as_current_span( + "ambient-invocation" + ) + if ambient + else nullcontext() + ) + try: + with scope: + invocation_context = context.get_current() + try: + return module.handler(event, lambda_context) + finally: + invocation_contexts.append( + (invocation_context, context.get_current()) + ) + finally: + assert context.get_current() is host_context + + try: + with DurableFunctionTestRunner(handler=entry, skip_time=False) as runner: + arn = runner.run_async( + input=json.dumps({"scenario": "user-function-context"}), + execution_timeout=30, + ) + result = runner.wait_for_result(arn, timeout=30) + history = runner.get_execution_history(arn) + assert result.status.value == "SUCCEEDED", result.error + assert result.result is not None + assert json.loads(result.result) == "context-complete" + assert [ + event.event_type + for event in history.events + if event.event_type.startswith("Wait") + ] == ["WaitStarted", "WaitSucceeded"] + + spans = exporter.get_finished_spans() + invocations = [span for span in spans if span.name == "Invocation"] + assert len(invocations) == 2 + assert [ + span.attributes["durable.invocation.status"] + for span in invocations + if span.attributes is not None + ] == ["PENDING", "SUCCEEDED"] + assert [ + span.attributes["durable.invocation.first"] + for span in invocations + if span.attributes is not None + ] == [True, False] + workflows = [span for span in spans if span.name == "Workflow"] + assert len(workflows) == 1 + assert workflows[0].context is not None + canonical_trace_id = workflows[0].context.trace_id + assert { + span.context.trace_id + for span in spans + if span.context is not None + and span.attributes is not None + and span.attributes.get("durable.execution.arn") == arn + } == {canonical_trace_id} + + probes = [span for span in spans if span.name.startswith("conformance.")] + # Count raw exports so replayed callback bodies and duplicate exports + # fail even if they reuse a span ID. + assert Counter(span.name for span in probes) == Counter( + { + f"conformance.{label}": count + for label, count in ( + {**dict.fromkeys(CALLBACK_PARENTS, 1), **HANDLER_COUNTS} + ).items() + } + ) + spans_by_id = { + (span.context.trace_id, span.context.span_id): span + for span in spans + if span.context is not None + } + handler_parent = ( + "Workflow" if plugin_type is ExecutionOtelPlugin else "Invocation" + ) + for probe_span in probes: + assert probe_span.context is not None + assert probe_span.context.trace_id == canonical_trace_id + assert probe_span.parent is not None + assert probe_span.parent.trace_id == canonical_trace_id + assert probe_span.attributes is not None + assert "durable.execution.arn" not in probe_span.attributes + label = str(probe_span.attributes["conformance.callback"]) + assert probe_span.name == f"conformance.{label}" + parent = spans_by_id[(canonical_trace_id, probe_span.parent.span_id)] + assert parent.name == CALLBACK_PARENTS.get(label, handler_parent) + assert parent.attributes is not None + assert parent.attributes["durable.execution.arn"] == arn + + for index, invocation in enumerate(invocations): + assert invocation.start_time is not None and invocation.end_time is not None + invocation_probes = [ + span + for span in probes + if span.start_time is not None + and invocation.start_time <= span.start_time <= invocation.end_time + ] + expected = ( + { + **dict.fromkeys(CALLBACK_PARENTS, 1), + "handler": 1, + "handler-restored": 1, + } + if index == 0 + else dict.fromkeys(HANDLER_COUNTS, 1) + ) + assert Counter(span.name for span in invocation_probes) == Counter( + {f"conformance.{label}": count for label, count in expected.items()} + ) + handler_probes = [ + span + for span in invocation_probes + if span.attributes is not None + and span.attributes["conformance.callback"] in HANDLER_COUNTS + ] + expected_parent = ( + workflows[0] if plugin_type is ExecutionOtelPlugin else invocation + ) + assert all( + span.parent == expected_parent.context for span in handler_probes + ) + + ambient_spans = [span for span in spans if span.name == "ambient-invocation"] + assert len(ambient_spans) == (2 if ambient else 0) + for span in ambient_spans: + assert span.context is not None + assert span.context.trace_id != canonical_trace_id + # Observe restoration on the runner's invocation threads, including + # PENDING, rather than checking only the pytest caller's context. + assert len(invocation_contexts) == 2 + assert all(before is after for before, after in invocation_contexts) + assert context.get_current() is before_context + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests/handlers/plugin/plugin_wait_replay_flag.py b/packages/aws-durable-execution-sdk-python-conformance-tests/handlers/plugin/plugin_wait_replay_flag.py index eef6d3eed..2998740d8 100644 --- a/packages/aws-durable-execution-sdk-python-conformance-tests/handlers/plugin/plugin_wait_replay_flag.py +++ b/packages/aws-durable-execution-sdk-python-conformance-tests/handlers/plugin/plugin_wait_replay_flag.py @@ -18,6 +18,7 @@ """ import json +from threading import Lock from typing import Any from aws_durable_execution_sdk_python.config import Duration, ParallelConfig @@ -31,13 +32,19 @@ ) +_log_lock = Lock() + + def _emit(record: dict[str, Any], execution_arn: str | None) -> None: # Prefix every plugin record with the execution ARN as a top-level field so # the conformance runner's CloudWatch JSON filter can scope logs to a single # execution. Omit the field when the ARN is unset (never invent a value). if execution_arn: record = {"durableExecutionArn": execution_arn, **record} - print(json.dumps(record), flush=True) + # Start and end hooks can run on different threads. Keep print's separate + # body/newline writes together so the runner receives one JSON per line. + with _log_lock: + print(json.dumps(record), flush=True) class WaitReplayFlagPlugin(DurableInstrumentationPlugin): diff --git a/packages/aws-durable-execution-sdk-python-conformance-tests/tests/test_plugin_wait_replay_flag.py b/packages/aws-durable-execution-sdk-python-conformance-tests/tests/test_plugin_wait_replay_flag.py new file mode 100644 index 000000000..0c6006043 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-conformance-tests/tests/test_plugin_wait_replay_flag.py @@ -0,0 +1,72 @@ +"""Concurrent plugin callbacks must emit separate parseable JSON records.""" + +from __future__ import annotations + +import importlib.util +import json +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from threading import Event +from typing import Any + +import pytest + +import aws_durable_execution_sdk_python.execution as execution + + +def test_concurrent_wait_hooks_keep_complete_stdout_records( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Import the real fixture without constructing a Lambda client: this test + # exercises its stdout producer, not a deployed durable invocation. + monkeypatch.setattr(execution, "durable_execution", lambda **_: lambda fn: fn) + path = ( + Path(__file__).resolve().parents[1] + / "handlers/plugin/plugin_wait_replay_flag.py" + ) + spec = importlib.util.spec_from_file_location("wait_replay_log_fixture", path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + + first_body = Event() + second_ready = Event() + second_done = Event() + chunks: list[str] = [] + + class FragmentingStdout: + def write(self, text: str) -> int: + chunks.append(text) + if '"operation-start"' in text: + first_body.set() + # print writes its body and newline separately. Permit the + # other real hook to run between them unless _emit serializes it. + second_done.wait(0.2) + return len(text) + + def flush(self) -> None: + pass + + records: list[dict[str, Any]] = [ + {"plugin": "CONFPLUGIN", "hook": "operation-start", "name": "long"}, + {"plugin": "CONFPLUGIN", "hook": "operation-end", "name": "short"}, + ] + + def emit_end() -> None: + second_ready.set() + assert first_body.wait(2) + module._emit(records[1], "execution-arn") + second_done.set() + + with monkeypatch.context() as capture: + capture.setattr("sys.stdout", FragmentingStdout()) + with ThreadPoolExecutor(max_workers=2) as executor: + end = executor.submit(emit_end) + assert second_ready.wait(2) + start = executor.submit(module._emit, records[0], "execution-arn") + start.result(timeout=2) + end.result(timeout=2) + actual = [json.loads(line) for line in "".join(chunks).splitlines() if line] + assert actual == [ + {"durableExecutionArn": "execution-arn", **record} for record in records + ] diff --git a/packages/aws-durable-execution-sdk-python-otel/README.md b/packages/aws-durable-execution-sdk-python-otel/README.md index 2c7e1bb1b..f1c4a0ff8 100644 --- a/packages/aws-durable-execution-sdk-python-otel/README.md +++ b/packages/aws-durable-execution-sdk-python-otel/README.md @@ -55,9 +55,11 @@ DURABLE_EXECUTION_PLUGINS=otel-execution cold start, so the handler does not need to import or explicitly register the plugin. -Automatic mutual-exclusion validation requires core SDK 2.1.0 or later together -with OTel 1.1.0 or later. OTel 1.1 remains compatible with core 2.0.x for existing -valid registrations; those older cores do not enforce the new group metadata. +OTel 1.1.0 requires core SDK 2.1.0 or later for the invocation-worker lifecycle +and automatic mutual-exclusion validation. Publish the redesigned core first, +then the plugin. Core 2.0.x runs invocation hooks on the caller and cannot provide +the new plugin's handler propagation and host-context isolation; that version +pair is not supported by OTel 1.1.0. Configure only one OTel view on every core version. `InvocationOtelPlugin` shows work within each Lambda invocation; `ExecutionOtelPlugin` shows logical operations across the whole execution. They emit overlapping telemetry and @@ -180,6 +182,38 @@ lambda_.Function( ) ``` +### Handler context propagation + +The core runs the existing Start hooks, handler, output preparation, resource +cleanup and End hooks on one invocation worker. Start finishes before checkpoint +processing begins. End follows registered-branch joins and checkpoint shutdown, +and output serialization or checkpoint errors retain their normal classification. +There is no separate handler-context plugin API. + +During Start, Invocation view preserves a valid active span on the canonical +execution trace; when it is absent or unrelated, it attaches the Invocation span. +Execution view attaches the Workflow span, including when an incoming span is on +the same trace. The Invocation span's own ambient parenting is separate from the +active context supplied to handler instrumentation. Both views preserve baggage. + +Start hooks run in registration order. A later successful plugin that deliberately +sets or clears the active span wins; OTel does not apply a second correction pass. +Place OTel after a span-replacing plugin when OTel's view-specific context is +desired. Baggage-only plugins that extend the current context can appear on either +side. The SDK also copies the coordinator's current bindings for each `map` or +`parallel` branch admission/resume when plugins are registered. Baggage and +other successful bindings reach these SDK-managed callbacks, while a branch's +changes cannot leak into siblings, the coordinator or a reused worker. +End hooks also retain registration order and reset tokens in their owning +Context; they do not promise a reverse-stack observation of other plugins' spans. +Existing registration, factory lifetime and checkpoint formats are unchanged. + +Upgrade both core to 2.1+ and OTel to 1.1+ for these guarantees. The released +OTel 1.0 plugin can run on the new core with worker/host isolation, but its +Invocation view does not attach the new fallback; a core-only upgrade does not +supply that plugin behavior. CI tests the new pair as installed wheels and tests +the actual released plugin separately, including before core 2.1 is on PyPI. + ### 3. In your Lambda handler (index.py) ```python @@ -330,6 +364,24 @@ OTel `OK` only for `SUCCEEDED`, `ERROR` when error details are delivered, and `CANCELLED`, `TIMED_OUT`, and `STOPPED`. The original durable operation status remains in `durable.operation.status`. +When a concurrent branch's terminal completion arrives before replay has created +its parent span, Invocation view retains that completion until the actual parent +span is registered. The terminal segment is exported under that parent before +control returns to the branch's user code. No ancestor is invented and SDK +completion delivery is unchanged. If the parent never becomes active, normal +invocation cleanup and flushing still run before the missing parent is reported. + +### Invocation context isolation + +With core 2.1+, invocation hooks and the handler run on one worker in an +invocation-local Context initialized from the host's bindings. Successful Start +bindings are visible to later hooks, the handler and resource cleanup. The host's +bindings remain unchanged after return, even if a plugin fails during Start or +End. A failed Start's bindings are discarded for subsequent work, while End runs +in that Start's original Context to preserve token ownership. This isolates +context-variable bindings; it does not undo mutations to shared objects or +external side effects. The isolation applies when plugins are registered. + ### Log Correlation When `enrich_logger=True` (the default), the plugin installs a logging filter on @@ -430,7 +482,7 @@ setups. ## Requirements - Python >= 3.11 -- `aws-durable-execution-sdk-python` >= 2.0.0 (core >= 2.1.0 with OTel >= 1.1.0 for automatic view-exclusivity validation) +- `aws-durable-execution-sdk-python` >= 2.1.0 (release the redesigned core before OTel 1.1.0) - An ADOT/community OpenTelemetry Lambda layer, or the `standalone` extra ## License diff --git a/packages/aws-durable-execution-sdk-python-otel/pyproject.toml b/packages/aws-durable-execution-sdk-python-otel/pyproject.toml index cbc2f97a2..461718714 100644 --- a/packages/aws-durable-execution-sdk-python-otel/pyproject.toml +++ b/packages/aws-durable-execution-sdk-python-otel/pyproject.toml @@ -22,7 +22,7 @@ classifiers = [ "Programming Language :: Python :: Implementation :: PyPy", ] dependencies = [ - "aws-durable-execution-sdk-python>=2.0.0", + "aws-durable-execution-sdk-python>=2.1.0", ] [project.entry-points."aws_durable_execution.plugins"] diff --git a/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/invocation_plugin.py b/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/invocation_plugin.py index efaf3b965..9712a81a7 100644 --- a/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/invocation_plugin.py +++ b/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/invocation_plugin.py @@ -73,6 +73,7 @@ {InvocationStatus.SUCCEEDED, InvocationStatus.FAILED} ) _TIMESTAMP_STEP_NANOS = 1_000 +_INVOCATION_CONTEXT_KEY = "__invocation_context__" _SpanAttributes = dict[str, str | bool | int] @@ -155,6 +156,9 @@ def __init__(self, config: OtelPluginConfig | None = None) -> None: self._span_time_floor_ns: int | None = None # Maps operation ID (None for root) to the active span. self._operation_spans: dict[str | None, Span] = {} + # A sibling checkpoint can report completion before replay enters its + # parent context. Retain the real event until that parent span exists. + self._pending_operation_ends: dict[str, list[OperationEndInfo]] = {} # Replay state supplied by CONTEXT operation START hooks. Missing # entries identify checkpointless contexts such as FLAT branches. self._context_operation_replays: dict[str, bool] = {} @@ -299,12 +303,8 @@ def get_current_span_context(self) -> SpanContext | None: context this is the active context span (attached in on_user_function_start). Unrelated ambient spans are ignored so logs stay correlated to the durable execution trace. - 2. The invocation span from the plugin registry. This is the path used - for top-level handler code: the invocation span is never attached to - the worker thread's context, so the registry is the only way to - resolve it. It also covers code between top-level operations, where - detaching the operation scope restores a context with no durable - span. + 2. The invocation span from the plugin registry, including lifecycle + phases where another plugin changed the active context. Returns: A valid SpanContext, or None if no span is active. @@ -494,10 +494,19 @@ def _start_span( links=links, ) self._operation_spans[registry_key] = span + pending_ends = ( + self._pending_operation_ends.pop(registry_key, []) + if registry_key is not None + else [] + ) if operation_id is None: self._span_time_floor_ns = span_start_time logger.debug("Started OTel span: %s", span) + # Registration and dequeue share the lock, so an arriving completion + # either sees its parent or is drained here. Do not hold it in callbacks. + for pending_end in pending_ends: + self.on_operation_end(pending_end) return span def _end_span( @@ -593,6 +602,17 @@ def on_invocation_start(self, info: InvocationStartInfo) -> None: attributes=self._extract_attributes(info), ) + # Start and End share the invocation worker and token-owning Context. + # Retain a valid same-trace ambient parent; otherwise bind this invocation. + ambient = trace.get_current_span().get_span_context() + invocation_span = self._get_span(None) + if invocation_span is not None and ( + not ambient.is_valid or ambient.trace_id != self._execution_trace_id + ): + self._attach_context( + _INVOCATION_CONTEXT_KEY, trace.set_span_in_context(invocation_span) + ) + # Cover handlers installed after construction as well. if self._enrich_logger: install_log_filter(self) @@ -673,11 +693,16 @@ def on_invocation_end(self, info: InvocationEndInfo) -> None: self._reset_state() return + # User work and checkpoint cleanup have finished. Release Start bindings + # in the Context that created their tokens before closing the spans. + self._detach_remaining_contexts() + # Spans are registered parent-first, so close pending spans in reverse # order to keep every child contained within its parent. with self._operation_spans_lock: operation_ids = list(reversed(self._operation_spans)) incomplete_attempt_span_keys = set(self._incomplete_attempt_span_keys) + unresolved_parents = tuple(self._pending_operation_ends) for operation_id in operation_ids: if operation_id: span = self._get_span(operation_id) @@ -726,6 +751,12 @@ def on_invocation_end(self, info: InvocationEndInfo) -> None: # Flush before Lambda freeze if hasattr(self._provider, "force_flush"): self._provider.force_flush() + if unresolved_parents: + # Keep the previous missing-parent failure observable, but finish + # normal token/span cleanup and flushing before reporting it. + raise ValueError( + f"No parent span found for deferred operation ends: {unresolved_parents}" + ) def _reset_state(self) -> None: """Clear per-invocation state for warm Lambda environment reuse.""" @@ -740,6 +771,7 @@ def _reset_state(self) -> None: self._span_time_floor_ns = None with self._operation_spans_lock: self._operation_spans = {} + self._pending_operation_ends = {} self._context_operation_replays = {} self._incomplete_attempt_span_keys = set() self._tracing_enabled = False @@ -783,7 +815,15 @@ def on_operation_end(self, info: OperationEndInfo) -> None: # The operation started in a prior invocation. Create a new # correlated segment and link it to the deterministic logical # operation context shared across invocations. - parent_span = self._resolve_parent_span(info.parent_id) + with self._operation_spans_lock: + parent_span = self._operation_spans.get(info.parent_id) + if parent_span is None and info.parent_id is not None: + self._pending_operation_ends.setdefault(info.parent_id, []).append( + info + ) + return + if parent_span is None: + parent_span = self._resolve_parent_span(info.parent_id) attributes = self._extract_attributes(info) span = self._start_span( operation_id=info.operation_id, diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_branch_context_int.py b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_branch_context_int.py new file mode 100644 index 000000000..c2ae779af --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_branch_context_int.py @@ -0,0 +1,179 @@ +"""Public SDK branch workers inherit instrumented bindings without sharing them.""" + +from __future__ import annotations + +import contextvars +import json +import threading +from typing import Any + +import pytest +from aws_durable_execution_sdk_python import DurableContext, durable_execution +from aws_durable_execution_sdk_python.config import Duration, MapConfig, ParallelConfig +from aws_durable_execution_sdk_python.concurrency.models import BatchResult +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationStatus, +) +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from opentelemetry import baggage, context, trace +from opentelemetry.sdk.trace import TracerProvider + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, +) + + +@pytest.mark.parametrize( + ("view", "failed_start"), + [ + (None, False), + (InvocationOtelPlugin, False), + (ExecutionOtelPlugin, False), + (InvocationOtelPlugin, True), + (ExecutionOtelPlugin, True), + ], +) +@pytest.mark.parametrize("kind", ["map", "parallel"]) +@pytest.mark.parametrize("concurrency", [1, 2]) +@pytest.mark.parametrize("resume", [False, True]) +def test_public_branch_context_propagation_and_isolation( + monkeypatch: pytest.MonkeyPatch, + view: type[InvocationOtelPlugin] | type[ExecutionOtelPlugin] | None, + failed_start: bool, + caplog: pytest.LogCaptureFixture, + kind: str, + concurrency: int, + resume: bool, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + marker = contextvars.ContextVar("branch-invocation-marker", default="unset") + observations: list[tuple[int, str, Any, bool]] = [] + bodies: list[int] = [] + visits = {0: 0, 1: 0} + statuses: list[InvocationStatus] = [] + starts: list[int] = [] + completed_ends: list[int] = [] + lock = threading.Lock() + barrier = threading.Barrier(2) if concurrency == 2 else None + provider = TracerProvider() + host = context.get_current() + + class Bind(DurableInstrumentationPlugin): + def on_invocation_start(self, info: Any) -> None: + starts.append(threading.get_ident()) + self.token = marker.set("plugin") + self.bag = context.attach(baggage.set_baggage("tenant", "present")) + + def on_invocation_end(self, info: Any) -> None: + context.detach(self.bag) + marker.reset(self.token) + # Record completion only after both tokens were reset in their owner. + completed_ends.append(threading.get_ident()) + statuses.append(info.status) + + class Broken(DurableInstrumentationPlugin): + def on_invocation_start(self, _info: Any) -> None: + marker.set("partial") + context.attach(baggage.set_baggage("tenant", "partial")) + raise ValueError("expected failed Start") + + plugins: list[DurableInstrumentationPlugin] = [] + otel_plugin: InvocationOtelPlugin | ExecutionOtelPlugin | None = None + if view is not None: + plugins.append(Bind()) + if failed_start: + plugins.append(Broken()) + otel_plugin = view( + OtelPluginConfig( + tracer_provider=provider, + enrich_logger=False, + context_extractor=lambda _: None, + ) + ) + plugins.append(otel_plugin) + + def branch(child: DurableContext, index: int) -> str: + with lock: + visits[index] += 1 + first_visit = visits[index] == 1 + observations.append( + ( + index, + marker.get(), + baggage.get_baggage("tenant"), + trace.get_current_span().get_span_context().is_valid, + ) + ) + if view is not None: + # Deliberately leave a binding behind: another logical branch or + # resume on this pool thread must still start with its parent's copy. + marker.set(f"branch-{index}") + if barrier is not None and first_visit: + barrier.wait(timeout=10) + + def step(_step: Any) -> str: + with lock: + bodies.append(index) + return f"saved-{index}" + + saved = child.step(step, name="save") + return saved + + def handler(_event: Any, durable: DurableContext) -> list[str]: + expected = "unset" if view is None else "plugin" + assert marker.get() == expected + token = marker.set("handler" if view is None else "plugin") + try: + result: BatchResult[str] = ( + durable.map( + [0, 1], + lambda child, item, index, items: branch(child, index), + name="mapped", + config=MapConfig(max_concurrency=concurrency), + ) + if kind == "map" + else durable.parallel( + [lambda child: branch(child, 0), lambda child: branch(child, 1)], + name="parallel", + config=ParallelConfig(max_concurrency=concurrency), + ) + ) + assert marker.get() == ("handler" if view is None else "plugin") + if resume: + durable.wait(Duration.from_seconds(1), name="resume") + return result.get_results() + finally: + marker.reset(token) + + wrapped = durable_execution(handler, plugins=plugins) + try: + with DurableFunctionTestRunner(handler=wrapped) as runner: + result = runner.run(input="{}", timeout=30) + assert result.status.value == "SUCCEEDED" + assert json.loads(result.result) == ["saved-0", "saved-1"] + assert sorted(bodies) == [0, 1] + assert {item[0] for item in observations} == {0, 1} + expected = "unset" if view is None else "plugin" + assert all( + item[1:] + == (expected, None if view is None else "present", view is not None) + for item in observations + ) + # The final wait replays the completed batch without rerunning branches. + assert visits == {0: 1, 1: 1} + if view is not None and resume: + assert InvocationStatus.PENDING in statuses + assert starts == completed_ends + if view is not None: + assert starts and statuses[-1] is InvocationStatus.SUCCEEDED + assert otel_plugin is not None and otel_plugin._context_tokens == {} + assert not any( + r.exc_info and r.name == "opentelemetry.context" for r in caplog.records + ) + assert marker.get() == "unset" + assert context.get_current() == host + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_invocation_wait_resume_int.py b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_invocation_wait_resume_int.py index 635634938..fa7036920 100644 --- a/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_invocation_wait_resume_int.py +++ b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_invocation_wait_resume_int.py @@ -2,6 +2,7 @@ from __future__ import annotations +from contextlib import nullcontext from dataclasses import replace from datetime import UTC, datetime from typing import Any @@ -25,12 +26,16 @@ OperationType, StepDetails, ) +from aws_durable_execution_sdk_python.plugin import DurableInstrumentationPlugin from aws_durable_execution_sdk_python_otel.deterministic_id_generator import ( derive_workflow_span_id, ) from aws_durable_execution_sdk_python_otel.execution_plugin import ExecutionOtelPlugin from aws_durable_execution_sdk_python_otel.invocation_plugin import InvocationOtelPlugin from aws_durable_execution_sdk_python_otel.otel_plugin_config import OtelPluginConfig +from opentelemetry import context as otel_context +from opentelemetry import trace +from opentelemetry.propagators.aws.aws_xray_propagator import AwsXRayPropagator from opentelemetry.sdk.trace import TracerProvider from opentelemetry.sdk.trace.export import SimpleSpanProcessor from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter @@ -253,3 +258,242 @@ def handler_impl(_event: Any, context: DurableContext) -> str: assert completed_wait_span.end_time is not None assert after_resume.start_time is not None assert completed_wait_span.end_time <= after_resume.start_time + + +class InheritedInvocationPlugin(InvocationOtelPlugin): + pass + + +class InheritedExecutionPlugin(ExecutionOtelPlugin): + pass + + +@pytest.mark.parametrize( + ("plugin_type", "extra_context_plugin"), + [ + (view, extra) + for view in (InvocationOtelPlugin, ExecutionOtelPlugin) + for extra in (False, True, "same", "unrelated", "absent") + ] + + [(InheritedInvocationPlugin, False), (InheritedExecutionPlugin, "absent")], +) +@pytest.mark.parametrize("reverse_plugins", [False, True]) +@pytest.mark.parametrize("fail_after_resume", [False, True]) +@pytest.mark.parametrize("ambient_kind", ["same", "unrelated", "absent"]) +def test_handler_user_spans_inherit_context_across_resume_and_failure( + monkeypatch: pytest.MonkeyPatch, + plugin_type: type[InvocationOtelPlugin] | type[ExecutionOtelPlugin], + fail_after_resume: bool, + ambient_kind: str, + extra_context_plugin: bool | str, + reverse_plugins: bool, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + monkeypatch.setenv("_X_AMZN_TRACE_ID", XRAY_TRACE_HEADER) + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = plugin_type( + OtelPluginConfig(tracer_provider=provider, enrich_logger=False) + ) + tracer = provider.get_tracer("customer") + before_context = otel_context.get_current() + calls: list[str] = [] + + def user_span(name: str) -> None: + # Ordinary instrumentation: the SDK/caller supplies the active parent. + span = tracer.start_span(name) + span.end() + + def step_body(_step_context: Any) -> str: + calls.append("step") + user_span("step-user") + return "saved" + + def handler_body(_event: Any, context: DurableContext) -> str: + if extra_context_plugin: + assert baggage.get_baggage("customer") == "present" + user_span("handler-entry") + saved = context.step(step_body, name="before-wait") + user_span("handler-after-step") + context.wait(Duration.from_seconds(1), name="context-wait") + user_span("handler-after-resume") + if fail_after_resume: + raise ValueError("handler failed after resume") + return saved + + # Start hooks share the invocation worker, in registration order. A later + # successful plugin may deliberately replace or clear the active span. + from opentelemetry import baggage + + class BaggagePlugin(DurableInstrumentationPlugin): + token: Any = None + + def on_invocation_start(self, _info: Any) -> None: + current = baggage.set_baggage("customer", "present") + if isinstance(extra_context_plugin, str): + # A real third-party invocation hook can bind or clear a span. + # User functions below still use ordinary implicit parenting. + parent = trace.SpanContext( + trace_id=(XRAY_TRACE_ID if extra_context_plugin == "same" else 1), + span_id=0xCAFE, + is_remote=False, + trace_flags=trace.TraceFlags(1), + ) + current = trace.set_span_in_context( + trace.INVALID_SPAN + if extra_context_plugin == "absent" + else trace.NonRecordingSpan(parent), + current, + ) + self.token = otel_context.attach(current) + + def on_invocation_end(self, _info: Any) -> None: + otel_context.detach(self.token) + self.token = None + + plugins: list[DurableInstrumentationPlugin] = [plugin] + if extra_context_plugin: + plugins.append(BaggagePlugin()) + if reverse_plugins: + plugins.reverse() + handler = durable_execution(handler_body, plugins=plugins) + remote = AwsXRayPropagator().extract({"X-Amzn-Trace-Id": XRAY_TRACE_HEADER}) + assert trace.get_current_span(remote).get_span_context().trace_id == XRAY_TRACE_ID + initial_operations = [_execution_operation()] + checkpoint, operations = _checkpoint_store(initial_operations) + ambient_ids: list[int] = [] + host_context = remote if ambient_kind == "same" else otel_context.Context() + try: + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient" + ) as client_class: + client = Mock() + client.checkpoint = checkpoint + client_class.initialize_client.return_value = client + # Standard host instrumentation supplies a same-trace Lambda span. + host_scope = ( + tracer.start_as_current_span("lambda-first", context=host_context) + if ambient_kind != "absent" + else nullcontext() + ) + with host_scope: + host = trace.get_current_span() + ambient_ids.append(host.get_span_context().span_id) + first = handler(_event(initial_operations), _lambda_context()) + assert ( + trace.get_current_span().get_span_context() + == host.get_span_context() + ) + assert first["Status"] == InvocationStatus.PENDING.value + assert otel_context.get_current() == before_context + resumed_operations = [ + replace( + operation, + status=OperationStatus.SUCCEEDED, + end_timestamp=datetime.now(UTC), + ) + if operation.name == "context-wait" + else operation + for operation in operations.values() + ] + wait_id = next( + operation.operation_id + for operation in resumed_operations + if operation.name == "context-wait" + ) + checkpoint, _ = _checkpoint_store(resumed_operations) + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient" + ) as client_class: + client = Mock() + client.checkpoint = checkpoint + client_class.initialize_client.return_value = client + host_scope = ( + tracer.start_as_current_span("lambda-resume", context=host_context) + if ambient_kind != "absent" + else nullcontext() + ) + with host_scope: + host = trace.get_current_span() + ambient_ids.append(host.get_span_context().span_id) + resumed = handler( + _event(resumed_operations, updated_operation_ids=[wait_id]), + _lambda_context(), + ) + assert ( + trace.get_current_span().get_span_context() + == host.get_span_context() + ) + assert resumed["Status"] == ( + InvocationStatus.FAILED.value + if fail_after_resume + else InvocationStatus.SUCCEEDED.value + ) + assert calls == ["step"] + assert otel_context.get_current() == before_context + spans = exporter.get_finished_spans() + expected_parents: list[int | None] + if isinstance(extra_context_plugin, str) and not reverse_plugins: + # No second OTel correction pass: the later successful Start wins. + expected_parents = [ + None if extra_context_plugin == "absent" else 0xCAFE + ] * 2 + expected_trace_id = {"same": XRAY_TRACE_ID, "unrelated": 1, "absent": None}[ + extra_context_plugin + ] + elif issubclass(plugin_type, ExecutionOtelPlugin): + expected_parents = [derive_workflow_span_id(EXECUTION_ARN)] * 2 + expected_trace_id = XRAY_TRACE_ID + elif reverse_plugins and isinstance(extra_context_plugin, str): + expected_parents = ( + [0xCAFE] * 2 + if extra_context_plugin == "same" + else [ + span.context.span_id + for span in spans + if span.name == "Invocation" and span.context is not None + ] + ) + expected_trace_id = XRAY_TRACE_ID + else: + expected_parents = ( + [*ambient_ids] + if ambient_kind == "same" + else [ + span.context.span_id + for span in spans + if span.name == "Invocation" and span.context is not None + ] + ) + expected_trace_id = XRAY_TRACE_ID + for name in ("handler-entry", "handler-after-step"): + users = [span for span in spans if span.name == name] + assert len(users) == 2 + assert [span.parent.span_id if span.parent else None for span in users] == ( + expected_parents + ) + assert all(span.context is not None for span in users) + if expected_trace_id is None: + assert all( + span.parent is None and span.context.trace_id != XRAY_TRACE_ID + for span in users + ) + else: + assert all(span.context.trace_id == expected_trace_id for span in users) + after_resume = next( + span for span in spans if span.name == "handler-after-resume" + ) + assert ( + after_resume.parent.span_id if after_resume.parent else None + ) == expected_parents[1] + step_user = next(span for span in spans if span.name == "step-user") + assert step_user.parent is not None + assert any( + span.name == "before-wait attempt 1" + and span.context is not None + and span.context.span_id == step_user.parent.span_id + for span in spans + ) + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_parent_creation_order_int.py b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_parent_creation_order_int.py new file mode 100644 index 000000000..b998de4d3 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_parent_creation_order_int.py @@ -0,0 +1,193 @@ +"""A real completion can precede its replayed sibling's parent span.""" + +from __future__ import annotations + +import json +import threading +import time +from typing import Any + +import pytest +from aws_durable_execution_sdk_python import DurableContext, durable_execution +from aws_durable_execution_sdk_python.config import ( + CallbackConfig, + Duration, + ParallelConfig, +) +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationStartInfo, + OperationEndInfo, + OperationType, + UserFunctionStartInfo, +) +from aws_durable_execution_sdk_python.serdes import JsonSerDes +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from opentelemetry import trace +from opentelemetry.sdk.trace import ReadableSpan, TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from aws_durable_execution_sdk_python_otel import InvocationOtelPlugin, OtelPluginConfig + + +def _span_start(span: ReadableSpan) -> int: + assert span.start_time is not None + return span.start_time + + +def test_early_sibling_callback_exports_under_real_parent_before_user_observation( + monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + provider = TracerProvider() + exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = InvocationOtelPlugin( + OtelPluginConfig( + tracer_provider=provider, + enrich_logger=False, + context_extractor=lambda _: None, + ) + ) + second_start = threading.Event() + both_completed = threading.Event() + end_observed = threading.Event() + parent_entered = threading.Event() + generations = [0] + observed: list[tuple[int, int]] = [] + hook_order: list[str] = [] + + class Gate(DurableInstrumentationPlugin): + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + generations[0] += 1 + if generations[0] == 2: + second_start.set() + assert both_completed.wait(10) + + def on_user_function_start(self, info: UserFunctionStartInfo) -> None: + if generations[0] == 2 and info.name == "parallel-branch-1": + hook_order.append("parent-gated") + assert end_observed.wait(10) + parent_entered.set() + + class ObserveEnd(DurableInstrumentationPlugin): + def on_operation_end(self, info: OperationEndInfo) -> None: + if ( + generations[0] == 2 + and info.name == "target-1" + and info.operation_type is OperationType.CALLBACK + ): + # This observer follows OTel in the actual SDK dispatch order. + assert not parent_entered.is_set() + hook_order.append("early-completion-observed") + end_observed.set() + + def branch(child: DurableContext, index: int) -> Any: + value = child.create_callback( + name=f"target-{index}", config=CallbackConfig(serdes=JsonSerDes()) + ).result() + if index == 1: + finished = [ + span + for span in exporter.get_finished_spans() + if span.name == "target-1" + ] + observed.append( + (len(finished), trace.get_current_span().get_span_context().span_id) + ) + with provider.get_tracer("customer").start_as_current_span( + "after-target-1" + ): + pass + return value + + def handler(_event: Any, durable: DurableContext) -> list[Any]: + result = durable.parallel( + [lambda child: branch(child, 0), lambda child: branch(child, 1)], + name="parallel", + config=ParallelConfig(max_concurrency=2), + ) + durable.wait(Duration.from_seconds(1), name="later-one") + durable.wait(Duration.from_seconds(1), name="later-two") + return result.get_results() + + wrapped = durable_execution(handler, plugins=[Gate(), plugin, ObserveEnd()]) + try: + with DurableFunctionTestRunner( + handler=wrapped, poll_interval=0.01, skip_time=False, execution_timeout=25 + ) as runner: + arn = runner.run_async(input="{}") + deadline = time.monotonic() + 10 + callbacks = {} + while time.monotonic() < deadline: + history = runner.get_execution_history( + arn, include_execution_data=True + ).events + callbacks = { + event.name: event.callback_started_details.callback_id + for event in history + if event.event_type == "CallbackStarted" + } + if len(callbacks) == 2 and any( + event.event_type == "InvocationCompleted" for event in history + ): + break + time.sleep(0.01) + assert set(callbacks) == {"target-0", "target-1"} + runner.send_callback_success( + callbacks["target-0"], result=json.dumps("left").encode() + ) + assert second_start.wait(10) + # The second callback completes after the invocation input snapshot, + # so another branch's checkpoint returns this genuine new completion. + runner.send_callback_success( + callbacks["target-1"], result=json.dumps("right").encode() + ) + both_completed.set() + result = runner.wait_for_result(arn, timeout=15) + history = runner.get_execution_history( + arn, include_execution_data=True + ).events + assert result.status.value == "SUCCEEDED" + assert json.loads(result.result) == ["left", "right"] + assert end_observed.is_set() and parent_entered.is_set() + assert len(observed) == 1 and observed[0][0] == 2 + spans = exporter.get_finished_spans() + # Keep every span: Invocation view exports a nonterminal segment in + # invocation one and a distinct terminal continuation in invocation two. + # Neither segment may stand in for the other or reappear on later replays. + for index in (0, 1): + segments = sorted( + [span for span in spans if span.name == f"target-{index}"], + key=_span_start, + ) + assert len(segments) == 2 + initial, terminal = segments + assert initial.attributes is not None and terminal.attributes is not None + assert initial.attributes["durable.operation.status"] == "STARTED" + assert terminal.attributes["durable.operation.status"] == "SUCCEEDED" + assert initial.context.span_id != terminal.context.span_id + parents = sorted( + [span for span in spans if span.name == f"parallel-branch-{index}"], + key=_span_start, + ) + assert len(parents) == 2 + assert initial.parent is not None + assert initial.parent.span_id == parents[0].context.span_id + assert terminal.parent is not None + assert terminal.parent.span_id == parents[1].context.span_id + assert initial.end_time is not None and terminal.start_time is not None + assert initial.end_time <= terminal.start_time + if index == 1: + assert terminal.parent.span_id == observed[0][1] + (marker,) = [span for span in spans if span.name == "after-target-1"] + assert terminal.end_time is not None and marker.start_time is not None + assert terminal.end_time <= marker.start_time + assert sum(event.event_type == "InvocationCompleted" for event in history) == 4 + assert not [record for record in caplog.records if record.exc_info] + assert plugin._context_tokens == {} + finally: + both_completed.set() + end_observed.set() + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/test_execution_plugin.py b/packages/aws-durable-execution-sdk-python-otel/tests/test_execution_plugin.py index 12edcc0ae..214a3a89d 100644 --- a/packages/aws-durable-execution-sdk-python-otel/tests/test_execution_plugin.py +++ b/packages/aws-durable-execution-sdk-python-otel/tests/test_execution_plugin.py @@ -28,6 +28,7 @@ from opentelemetry import baggage, trace from opentelemetry.context import Context from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.sampling import ALWAYS_ON, ALWAYS_OFF, Sampler from opentelemetry.sdk.trace.export import SimpleSpanProcessor from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from opentelemetry.trace import ( @@ -1622,3 +1623,66 @@ def test_nested_suspension_unwinds_scopes_in_reverse_order(): plugin.on_invocation_end(_invocation_end_info()) assert plugin._context_tokens == {} + + +@pytest.mark.parametrize("ambient_kind", ["same", "unrelated", "absent"]) +@pytest.mark.parametrize("raises", [False, True]) +@pytest.mark.parametrize("sampler", [ALWAYS_ON, ALWAYS_OFF]) +def test_invocation_hooks_bind_parent_and_restore_baggage( + ambient_kind: str, raises: bool, sampler: Sampler +) -> None: + provider = TracerProvider(sampler=sampler) + plugin = ExecutionOtelPlugin( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + caller = baggage.set_baggage("tenant", "hook-test", Context()) + ambient = None + if ambient_kind != "absent": + ambient = SpanContext( + trace_id=_to_otel_trace_id(EXECUTION_ARN, START_TIME) + if ambient_kind == "same" + else 1, + span_id=0x42, + is_remote=False, + trace_flags=TraceFlags(1), + ) + caller = trace.set_span_in_context(NonRecordingSpan(ambient), caller) + token = otel_context.attach(caller) + error = ValueError("handler error") + try: + plugin.on_invocation_start(_invocation_start_info()) + expected = trace.get_current_span().get_span_context() + assert expected.trace_id == _to_otel_trace_id(EXECUTION_ARN, START_TIME) + assert expected.span_id == derive_workflow_span_id(EXECUTION_ARN) + + def body() -> None: + active = trace.get_current_span().get_span_context() + assert active == expected + assert active.is_valid + assert baggage.get_baggage("tenant") == "hook-test" + if raises: + raise error + + try: + if raises: + with pytest.raises(ValueError) as caught: + body() + assert caught.value is error + else: + body() + finally: + plugin.on_invocation_end( + _invocation_end_info( + InvocationStatus.FAILED if raises else InvocationStatus.SUCCEEDED + ) + ) + assert plugin._context_tokens == {} + assert otel_context.get_current() is caller + assert baggage.get_baggage("tenant") == "hook-test" + finally: + otel_context.detach(token) + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/test_invocation_plugin.py b/packages/aws-durable-execution-sdk-python-otel/tests/test_invocation_plugin.py index 1894d2b48..b1f863c68 100644 --- a/packages/aws-durable-execution-sdk-python-otel/tests/test_invocation_plugin.py +++ b/packages/aws-durable-execution-sdk-python-otel/tests/test_invocation_plugin.py @@ -724,6 +724,7 @@ def test_operation_end_without_start_links_previous_logical_operation(): assert ( span.attributes["durable.operation.status"] == OperationStatus.SUCCEEDED.value ) + plugin.on_invocation_end(_invocation_end_info()) def test_continuation_span_uses_current_start_and_end_times(): @@ -753,6 +754,7 @@ def test_continuation_span_uses_current_start_and_end_times(): span = exporter.get_finished_spans()[0] assert invocation_span.start_time <= span.start_time assert before_callback <= span.start_time <= span.end_time <= after_callback + plugin.on_invocation_end(_invocation_end_info()) def test_resume_operation_timestamps_do_not_precede_current_invocation(): @@ -812,6 +814,7 @@ def test_resume_operation_timestamps_do_not_precede_current_invocation(): assert invocation_span.start_time <= after_resume_span.start_time assert after_resume_span.parent is not None assert after_resume_span.parent.span_id == invocation_span.context.span_id + plugin.on_invocation_end(_invocation_end_info()) def test_ordered_timestamps_are_thread_safe(): @@ -878,6 +881,7 @@ def test_retried_operation_uses_fresh_id_and_links_previous_logical_operation(): derive_workflow_span_id(EXECUTION_ARN), operation_id_to_span_id(EXECUTION_ARN, operation_id), } + plugin.on_invocation_end(_invocation_end_info()) def test_step_operation_span_parents_attempt_span(): @@ -1028,6 +1032,7 @@ def test_user_function_callbacks_emit_attempt_span_attributes(): == UserFunctionOutcome.SUCCEEDED.value ) assert "durable.operation.status" not in span.attributes + plugin.on_invocation_end(_invocation_end_info()) def test_step_attempt_span_name_includes_attempt_number(): @@ -1070,6 +1075,7 @@ def test_step_attempt_span_name_includes_attempt_number(): span = exporter.get_finished_spans()[0] assert span.name == "fetch-user attempt 2" + plugin.on_invocation_end(_invocation_end_info()) def test_step_attempt_span_name_defaults_to_first_attempt(): @@ -1112,6 +1118,7 @@ def test_step_attempt_span_name_defaults_to_first_attempt(): span = exporter.get_finished_spans()[0] assert span.name == "fetch-user attempt 1" + plugin.on_invocation_end(_invocation_end_info()) @pytest.mark.parametrize( @@ -1213,6 +1220,7 @@ def test_context_span_waits_for_terminal_status_and_omits_attempt_attributes( assert "durable.attempt.number" not in span.attributes assert "durable.attempt.outcome" not in span.attributes assert span.status.status_code is expected_span_status + plugin.on_invocation_end(_invocation_end_info()) def test_span_registry_helpers_can_be_called_from_multiple_threads(): @@ -1257,8 +1265,9 @@ def test_user_function_end_restores_enclosing_context(): # After the step, the enclosing context is restored and no scope is left # behind. Log correlation resolves the invocation span from the registry. assert otel_context.get_current() == enclosing_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} assert plugin.get_current_span_context().span_id == invocation_span_id + plugin.on_invocation_end(_invocation_end_info()) def test_user_function_start_preserves_baggage_in_current_context(): @@ -1292,7 +1301,8 @@ def test_user_function_end_restores_enclosing_context_on_failure(): ) assert otel_context.get_current() == enclosing_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} + plugin.on_invocation_end(_invocation_end_info()) def test_user_function_end_restores_enclosing_context_across_multiple_steps(): @@ -1309,8 +1319,9 @@ def test_user_function_end_restores_enclosing_context_across_multiple_steps(): # Between each step the context is back to where it started, and log # correlation still resolves the invocation span. assert otel_context.get_current() == enclosing_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} assert plugin.get_current_span_context().span_id == invocation_span_id + plugin.on_invocation_end(_invocation_end_info()) # ---------------------------------------------------------------------- @@ -1332,6 +1343,7 @@ def test_get_current_span_context_returns_invocation_span_at_top_level(): invocation_span = plugin._get_span(None) assert span_context is not None assert span_context.span_id == invocation_span.get_span_context().span_id + plugin.on_invocation_end(_invocation_end_info()) def test_get_current_span_context_returns_operation_span_inside_step(): @@ -1363,6 +1375,7 @@ def test_get_current_span_context_returns_invocation_span_between_steps(): invocation_span = plugin._get_span(None) assert span_context is not None assert span_context.span_id == invocation_span.get_span_context().span_id + plugin.on_invocation_end(_invocation_end_info()) # ---------------------------------------------------------------------- @@ -1441,11 +1454,12 @@ def test_top_level_step_end_falls_back_to_invocation_for_correlation(): plugin.on_user_function_start(_user_function_start_info(operation_id)) plugin.on_user_function_end(_user_function_end_info(operation_id)) - # No durable span is attached at the top level, so the registry fallback - # supplies the invocation span for log correlation. + # The invocation Start binding is restored after the step, including + # between-step user instrumentation and log correlation. assert otel_context.get_current() == enclosing_context - assert not trace.get_current_span().get_span_context().is_valid + assert trace.get_current_span().get_span_context().span_id == invocation_span_id assert plugin.get_current_span_context().span_id == invocation_span_id + plugin.on_invocation_end(_invocation_end_info()) def test_get_current_span_context_returns_context_span_between_nested_steps(): @@ -1709,6 +1723,7 @@ def test_replayed_context_span_links_previous_logical_operation(): derive_workflow_span_id(EXECUTION_ARN), operation_id_to_span_id(EXECUTION_ARN, operation_id), } + plugin.on_invocation_end(_invocation_end_info()) def test_checkpointed_context_first_span_uses_deterministic_id(): @@ -1867,7 +1882,8 @@ def test_child_context_end_restores_context_active_before_it(): ) assert otel_context.get_current() == enclosing_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} + plugin.on_invocation_end(_invocation_end_info()) def test_nested_scopes_are_released_without_accumulating(): @@ -1891,13 +1907,14 @@ def test_nested_scopes_are_released_without_accumulating(): ) # The inner step restored the child-context scope, not a copy of it. assert otel_context.get_current() == inside_context - assert set(plugin._context_tokens) == {context_id} + assert set(plugin._context_tokens) == {context_id, "__invocation_context__"} plugin.on_user_function_end( _user_function_end_info(context_id, operation_type=OperationType.CONTEXT) ) assert otel_context.get_current() == before_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} + plugin.on_invocation_end(_invocation_end_info()) def test_invocation_end_releases_scope_of_suspended_user_function(): @@ -1981,7 +1998,7 @@ def test_detach_ignores_token_attached_on_another_thread(): plugin._detach_context("step-1:attempt:1") - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} assert otel_context.get_current() == before_context plugin.on_invocation_end(_invocation_end_info()) @@ -2072,7 +2089,7 @@ def test_reentered_step_attempt_releases_the_previous_scope(): plugin.on_user_function_end(_user_function_end_info(operation_id)) assert otel_context.get_current() == before_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} plugin.on_invocation_end(_invocation_end_info()) @@ -2166,7 +2183,7 @@ def test_nested_suspension_unwinds_scopes_in_reverse_order(): ) ) assert otel_context.get_current() == before_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} # Neither span is ended: both operations are still in flight. assert not exporter.get_finished_spans() @@ -2199,6 +2216,201 @@ def test_nested_suspension_unwinds_scopes_in_reverse_order(): _user_function_end_info("ctx-outer", operation_type=OperationType.CONTEXT) ) assert otel_context.get_current() == before_context - assert plugin._context_tokens == {} + assert set(plugin._context_tokens) == {"__invocation_context__"} plugin.on_invocation_end(_invocation_end_info()) + + +@pytest.mark.parametrize("ambient_kind", ["same", "unrelated", "absent"]) +@pytest.mark.parametrize("raises", [False, True]) +@pytest.mark.parametrize("sampler", [ALWAYS_ON, ALWAYS_OFF]) +def test_invocation_hooks_bind_parent_and_restore_baggage( + ambient_kind: str, raises: bool, sampler: Sampler +) -> None: + provider = TracerProvider(sampler=sampler) + plugin = InvocationOtelPlugin( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + caller = baggage.set_baggage("tenant", "hook-test", Context()) + ambient = None + if ambient_kind != "absent": + ambient = SpanContext( + trace_id=_to_otel_trace_id(EXECUTION_ARN, START_TIME) + if ambient_kind == "same" + else 1, + span_id=0x42, + is_remote=False, + trace_flags=TraceFlags(1), + ) + caller = trace.set_span_in_context(NonRecordingSpan(ambient), caller) + token = otel_context.attach(caller) + error = ValueError("handler error") + try: + plugin.on_invocation_start(_invocation_start_info()) + expected = ( + ambient if ambient_kind == "same" else plugin.get_current_span_context() + ) + + def body() -> None: + active = trace.get_current_span().get_span_context() + assert active == expected + assert active.is_valid + assert baggage.get_baggage("tenant") == "hook-test" + if raises: + raise error + + try: + if raises: + with pytest.raises(ValueError) as caught: + body() + assert caught.value is error + else: + body() + finally: + plugin.on_invocation_end( + _invocation_end_info( + InvocationStatus.FAILED if raises else InvocationStatus.SUCCEEDED + ) + ) + assert plugin._context_tokens == {} + assert otel_context.get_current() is caller + assert baggage.get_baggage("tenant") == "hook-test" + finally: + otel_context.detach(token) + provider.shutdown() + + +@pytest.mark.parametrize( + ("status", "error", "expected"), + [ + (OperationStatus.SUCCEEDED, None, StatusCode.OK), + (OperationStatus.FAILED, None, StatusCode.UNSET), + (OperationStatus.CANCELLED, None, StatusCode.UNSET), + (OperationStatus.TIMED_OUT, None, StatusCode.UNSET), + (OperationStatus.STOPPED, None, StatusCode.UNSET), + ( + OperationStatus.FAILED, + ErrorObject(message="failure", type="Example", data=None, stack_trace=None), + StatusCode.ERROR, + ), + ], +) +def test_deferred_end_retains_real_parent_and_status_mapping( + status: OperationStatus, error: ErrorObject | None, expected: StatusCode +) -> None: + plugin, exporter = _create_plugin() + plugin.on_invocation_start(_invocation_start_info()) + end = OperationEndInfo( + operation_id="child", + operation_type=OperationType.WAIT, + sub_type=None, + name="deferred-child", + parent_id="parent", + start_time=START_TIME, + is_replayed=False, + status=status, + end_time=END_TIME, + error=error, + ) + plugin.on_operation_end(end) + assert not exporter.get_finished_spans() + assert plugin._pending_operation_ends == {"parent": [end]} + plugin.on_user_function_start( + _user_function_start_info("parent", operation_type=OperationType.CONTEXT) + ) + parent = plugin._get_span("parent") + assert parent is not None + (child,) = exporter.get_finished_spans() + assert child.parent is not None and child.attributes is not None + assert child.parent.span_id == parent.get_span_context().span_id + assert child.status.status_code is expected + assert child.attributes["durable.operation.status"] == status.value + assert len(child.events) == (1 if error is not None else 0) + assert plugin._pending_operation_ends == {} + plugin.on_user_function_end( + _user_function_end_info("parent", operation_type=OperationType.CONTEXT) + ) + plugin.on_invocation_end(_invocation_end_info()) + + +def test_deferred_descendants_drain_from_real_parent_completion_in_child_first_order() -> ( + None +): + plugin, exporter = _create_plugin() + plugin.on_invocation_start(_invocation_start_info()) + for operation_id, parent_id, kind in [ + ("child", "inner", OperationType.WAIT), + ("inner", "outer", OperationType.CONTEXT), + ]: + plugin.on_operation_end( + OperationEndInfo( + operation_id=operation_id, + operation_type=kind, + sub_type=None, + name=operation_id, + parent_id=parent_id, + start_time=START_TIME, + is_replayed=False, + status=OperationStatus.SUCCEEDED, + end_time=END_TIME, + ) + ) + assert not exporter.get_finished_spans() + plugin.on_user_function_start( + _user_function_start_info("outer", operation_type=OperationType.CONTEXT) + ) + child, inner = exporter.get_finished_spans() + assert [child.name, inner.name] == ["child", "inner"] + assert child.parent is not None and inner.parent is not None + assert child.parent.span_id == inner.context.span_id + outer = plugin._get_span("outer") + assert outer is not None + assert inner.parent.span_id == outer.get_span_context().span_id + assert child.end_time is not None and inner.end_time is not None + assert child.end_time <= inner.end_time + assert plugin._pending_operation_ends == {} + plugin.on_user_function_end( + _user_function_end_info("outer", operation_type=OperationType.CONTEXT) + ) + plugin.on_invocation_end(_invocation_end_info()) + + +def test_unresolved_parent_is_reported_after_cleanup_and_not_carried_into_reuse() -> ( + None +): + plugin, exporter = _create_plugin() + before = otel_context.get_current() + plugin.on_invocation_start(_invocation_start_info()) + plugin.on_operation_end( + OperationEndInfo( + operation_id="missing-child", + operation_type=OperationType.WAIT, + sub_type=None, + name="missing-child", + parent_id="missing-parent", + start_time=START_TIME, + is_replayed=False, + status=OperationStatus.SUCCEEDED, + end_time=END_TIME, + ) + ) + with pytest.raises(ValueError, match="No parent span found for deferred"): + plugin.on_invocation_end(_invocation_end_info()) + assert otel_context.get_current() == before + assert plugin._operation_spans == {} + assert plugin._pending_operation_ends == {} + assert plugin._context_tokens == {} + assert {span.name for span in exporter.get_finished_spans()} == { + "Invocation", + "Workflow", + } + plugin.on_invocation_start(_invocation_start_info()) + plugin.on_invocation_end(_invocation_end_info(InvocationStatus.PENDING)) + assert not [ + span for span in exporter.get_finished_spans() if span.name == "missing-child" + ] + assert plugin._pending_operation_ends == {} diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/test_package_metadata.py b/packages/aws-durable-execution-sdk-python-otel/tests/test_package_metadata.py index ac9af3939..b7a34ebda 100644 --- a/packages/aws-durable-execution-sdk-python-otel/tests/test_package_metadata.py +++ b/packages/aws-durable-execution-sdk-python-otel/tests/test_package_metadata.py @@ -6,8 +6,7 @@ PACKAGE_ROOT = Path(__file__).resolve().parents[1] REPOSITORY_ROOT = PACKAGE_ROOT.parents[1] -CORE_DEPENDENCY = "aws-durable-execution-sdk-python>=2.0.0" -EXCLUSIVITY_TEST_CORE = "aws-durable-execution-sdk-python>=2.1.0" +CORE_DEPENDENCY = "aws-durable-execution-sdk-python>=2.1.0" TEST_OTEL_DEPENDENCIES = { "opentelemetry-sdk>=1.20.0", "opentelemetry-propagator-aws-xray", @@ -68,7 +67,7 @@ def test_test_environments_install_layer_provided_dependencies() -> None: "test", "dev-otel", "dev-examples", - "test-pypi-otel", + "test-wheel-otel", "test-pypi-examples", ): assert TEST_OTEL_DEPENDENCIES <= set( @@ -77,22 +76,7 @@ def test_test_environments_install_layer_provided_dependencies() -> None: assert TEST_OTEL_DEPENDENCIES <= set(environments["types"]["extra-dependencies"]) -def test_pypi_compatibility_environment_uses_compatible_core_sdk() -> None: - dependencies = _load_pyproject(REPOSITORY_ROOT / "pyproject.toml")["tool"]["hatch"][ - "envs" - ]["test-pypi-otel"]["dependencies"] - - assert EXCLUSIVITY_TEST_CORE in dependencies - - -def test_pypi_otel_environment_installs_lifecycle_test_runner() -> None: - dependencies = _load_pyproject(REPOSITORY_ROOT / "pyproject.toml")["tool"]["hatch"][ - "envs" - ]["test-pypi-otel"]["dependencies"] - assert "aws-durable-execution-sdk-python-testing>=1.2.1" in dependencies - - -def test_core_dependency_preserves_previously_supported_releases() -> None: +def test_core_dependency_requires_worker_lifecycle_release() -> None: dependencies = _load_pyproject(PACKAGE_ROOT / "pyproject.toml")["project"][ "dependencies" ] @@ -101,32 +85,23 @@ def test_core_dependency_preserves_previously_supported_releases() -> None: for value in dependencies if Requirement(value).name == "aws-durable-execution-sdk-python" ) - assert requirement.specifier.contains("2.0.0") - assert requirement.specifier.contains("2.0.1") + assert not requirement.specifier.contains("2.0.0") + assert not requirement.specifier.contains("2.0.1") assert requirement.specifier.contains("2.1.0") -def test_pypi_otel_environment_does_not_shadow_installed_core() -> None: - environment = _load_pyproject(REPOSITORY_ROOT / "pyproject.toml")["tool"]["hatch"][ +def test_wheel_lanes_do_not_shadow_installed_artifacts() -> None: + environments = _load_pyproject(REPOSITORY_ROOT / "pyproject.toml")["tool"]["hatch"][ "envs" - ]["test-pypi-otel"] - assert environment["workspace"]["members"] == [ - "packages/aws-durable-execution-sdk-python-otel" ] - - -def test_legacy_lane_retains_supported_core_20_range() -> None: - environment = _load_pyproject(REPOSITORY_ROOT / "pyproject.toml")["tool"]["hatch"][ - "envs" - ]["test-pypi-otel-legacy"] - requirement = next( - Requirement(value) - for value in environment["dependencies"] - if Requirement(value).name == "aws-durable-execution-sdk-python" + for name in ("test-wheel-otel", "test-wheel-otel-legacy"): + assert environments[name]["workspace"]["members"] == [] + assert environments[name]["detached"] is True + assert ( + "aws-durable-execution-sdk-python-otel==1.0.0" + in environments["test-wheel-otel-legacy"]["dependencies"] + ) + assert ( + "aws-durable-execution-sdk-python-testing>=1.2.1" + in environments["test-wheel-otel"]["dependencies"] ) - assert requirement.specifier.contains("2.0.0") - assert requirement.specifier.contains("2.0.1") - assert not requirement.specifier.contains("2.1.0") - assert environment["workspace"]["members"] == [ - "packages/aws-durable-execution-sdk-python-otel" - ] diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py index 82880a1ba..2ab02fdce 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py @@ -442,7 +442,8 @@ def mark_state_delivered(self) -> None: delivers state in the invocation input. So the list is reset when that input is built, not when the invocation completes: an operation that completes while the handler is still running is - reported on the next invocation. + reported on the next invocation unless a checkpoint response has + already delivered it through ``advance_handler_seen``. """ self.updated_operation_ids = [] @@ -647,6 +648,7 @@ def complete_callback_success( end_timestamp=now if now is not None else real_now(), callback_details=updated_callback_details, ) + self._record_updated_operation(operation.operation_id) return self.operations[index] def complete_callback_failure( @@ -663,8 +665,11 @@ def complete_callback_failure( self.touch_operation(operation.operation_id) updated_callback_details = None if operation.callback_details: + # Match CallbackDetails.from_dict without depending on a store + # serialization round trip: an empty wire Error has no details. updated_callback_details = replace( - operation.callback_details, error=error + operation.callback_details, + error=error if error is not None and error.to_dict() else None, ) self.operations[index] = replace( @@ -673,6 +678,7 @@ def complete_callback_failure( end_timestamp=now if now is not None else real_now(), callback_details=updated_callback_details, ) + self._record_updated_operation(operation.operation_id) return self.operations[index] def complete_callback_timeout( @@ -699,6 +705,7 @@ def complete_callback_timeout( end_timestamp=now if now is not None else real_now(), callback_details=updated_callback_details, ) + self._record_updated_operation(operation.operation_id) return self.operations[index] def complete_chained_invoke( @@ -880,6 +887,13 @@ def advance_handler_seen(self, seq: int) -> None: smaller or equal values are ignored.""" if seq > self.execution.handler_seen_seq: self.execution.handler_seen_seq = seq + # A checkpoint has delivered these updates to the running handler. + # Retain only changes newer than that response for the next input. + self.execution.updated_operation_ids = [ + operation_id + for operation_id in self.execution.updated_operation_ids + if self.execution.operation_last_touched_seq.get(operation_id, 0) > seq + ] # --- internals ------------------------------------------------- diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py index 43026ec60..fdccb2cff 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py @@ -677,7 +677,8 @@ class EventError: @classmethod def from_dict(cls, data: dict) -> EventError: payload = None - if payload_data := data.get("Payload"): + payload_data = data.get("Payload") + if payload_data is not None: payload = ErrorObject.from_dict(payload_data) return cls( @@ -2239,6 +2240,14 @@ def create_callback_event_failed(cls, context: EventCreationContext) -> Event: event_error: EventError | None = ( EventError.from_details(callback_details) if callback_details else None ) + if ( + context.include_execution_data + and callback_details is not None + and callback_details.error is None + ): + # Detailed service history retains an empty Error.Payload object. + # This projection must not turn the SDK-facing absent error into one. + event_error = EventError(payload=ErrorObject.from_dict({}), truncated=False) return cls( event_type=EventType.CALLBACK_FAILED.value, event_timestamp=context.end_timestamp, @@ -2744,7 +2753,9 @@ def events_to_operations(events: list[Event]) -> list[Operation]: callback_details=CallbackDetails( callback_id=callback_id, result=result, - error=error, + # History preserves a present empty Payload object, while + # CallbackDetails uses None for an empty wire error. + error=error if error is not None and error.to_dict() else None, ), ) diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/e2e/callback_updated_operations_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/e2e/callback_updated_operations_test.py new file mode 100644 index 000000000..cae94c294 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/tests/e2e/callback_updated_operations_test.py @@ -0,0 +1,253 @@ +"""Real callback completions are delivered once before resumed user code.""" + +from __future__ import annotations + +import json +import time +from pathlib import Path +from queue import Queue +from threading import Event +from typing import Any + +import pytest +from aws_durable_execution_sdk_python import ( + DurableContext, + StepContext, + durable_execution, + durable_step, +) +from aws_durable_execution_sdk_python.config import ( + CallbackConfig, + Duration, + WaitForCallbackConfig, +) +from aws_durable_execution_sdk_python.exceptions import CallbackError +from aws_durable_execution_sdk_python.lambda_service import ErrorObject +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationStartInfo, + OperationEndInfo, + UserFunctionStartInfo, +) +from aws_durable_execution_sdk_python.serdes import JsonSerDes +from aws_durable_execution_sdk_python.types import WaitForCallbackContext + +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner +from aws_durable_execution_sdk_python_testing.stores.filesystem import ( + FileSystemExecutionStore, +) + + +class CallbackObserver(DurableInstrumentationPlugin): + def __init__(self) -> None: + self.invocations: list[InvocationStartInfo] = [] + self.ends: list[OperationEndInfo] = [] + self.order: list[str] = [] + + def on_invocation_start(self, info: InvocationStartInfo) -> None: + self.invocations.append(info) + + def on_operation_end(self, info: OperationEndInfo) -> None: + if info.name == "target": + self.ends.append(info) + self.order.append("target-end") + + def on_user_function_start(self, info: UserFunctionStartInfo) -> None: + if info.name == "observed": + self.order.append("observed-start") + + +def _submit(_callback_id: str, _context: WaitForCallbackContext) -> None: + return None + + +def _suspended_callback(runner: DurableFunctionTestRunner, arn: str, name: str) -> str: + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + events = runner.get_execution_history(arn, include_execution_data=True).events + starts = [ + event + for event in events + if event.event_type == "CallbackStarted" and event.name == name + ] + if starts and any( + event.event_type == "InvocationCompleted" + and event.event_id > starts[0].event_id + for event in events + ): + details = starts[0].callback_started_details + assert details is not None + assert details.callback_id is not None + return details.callback_id + time.sleep(0.01) + raise AssertionError(f"Callback {name} did not suspend") + + +@pytest.mark.parametrize("outcome", ["success", "failure", "timeout"]) +@pytest.mark.parametrize("filesystem", [False, True]) +def test_callback_update_is_consumed_before_user_code( + outcome: str, filesystem: bool, tmp_path: Path +) -> None: + observer = CallbackObserver() + marker_calls: list[str] = [] + + @durable_step + def observed(_context: StepContext, value: str) -> str: + marker_calls.append(value) + return value + + def handler(_event: Any, context: DurableContext) -> str: + config = CallbackConfig( + timeout=Duration.from_seconds(1 if outcome == "timeout" else 30), + serdes=JsonSerDes(), + ) + try: + target = context.create_callback(name="target", config=config).result() + except CallbackError: + target = outcome + assert isinstance(target, str) + saved = context.step(observed(target), name="observed") + callback_config = WaitForCallbackConfig(serdes=JsonSerDes()) + one = context.wait_for_callback(_submit, name="one", config=callback_config) + two = context.wait_for_callback(_submit, name="two", config=callback_config) + return "/".join((saved, one, two)) + + wrapped = durable_execution(handler, plugins=[observer]) + store = FileSystemExecutionStore.create(tmp_path) if filesystem else None + with DurableFunctionTestRunner( + handler=wrapped, + store=store, + skip_time=False, + poll_interval=0.01, + execution_timeout=25, + ) as runner: + arn = runner.run_async(input="{}") + callback_id = _suspended_callback(runner, arn, "target") + if outcome == "success": + runner.send_callback_success( + callback_id, result=json.dumps("target").encode() + ) + elif outcome == "failure": + runner.send_callback_failure( + callback_id, error=ErrorObject.from_message("explicit callback failure") + ) + # The timeout case uses the runner's real scheduled callback deadline. + for name in ["one", "two"]: + callback_id = _suspended_callback(runner, arn, name + " create callback id") + runner.send_callback_success(callback_id, result=json.dumps(name).encode()) + result = runner.wait_for_result(arn, timeout=10) + + expected = "target" if outcome == "success" else outcome + assert result.status.value == "SUCCEEDED" + assert result.result is not None + assert json.loads(result.result) == expected + "/one/two" + assert marker_calls == [expected] + assert len(observer.invocations) == 4 + assert len(observer.ends) == 1 + terminal = observer.ends[0] + assert ( + terminal.status.value + == {"success": "SUCCEEDED", "failure": "FAILED", "timeout": "TIMED_OUT"}[ + outcome + ] + ) + assert terminal.is_replayed is False + assert set(observer.invocations[1].updated_operations) == {terminal.operation_id} + assert all( + terminal.operation_id not in invocation.updated_operations + for invocation in observer.invocations[2:] + ) + assert observer.order == ["target-end", "observed-start"] + if outcome == "failure": + assert terminal.error is not None + assert terminal.error.message == "explicit callback failure" + + +@pytest.mark.parametrize("outcome", ["success", "failure", "timeout"]) +@pytest.mark.parametrize("filesystem", [False, True]) +def test_callback_delivered_by_checkpoint_is_not_updated_on_later_resume( + outcome: str, filesystem: bool, tmp_path: Path +) -> None: + observer = CallbackObserver() + submitter_entered = Event() + release_submitter = Event() + submission_calls: list[str] = [] + published_callbacks: Queue[str] = Queue() + + @durable_step + def submit(_context: StepContext, callback_id: str) -> str: + submission_calls.append(callback_id) + published_callbacks.put(callback_id) + submitter_entered.set() + assert release_submitter.wait(10) + return "submitted" + + def handler(_event: Any, context: DurableContext) -> str: + callback = context.create_callback( + name="target", + config=CallbackConfig( + timeout=Duration.from_seconds(1 if outcome == "timeout" else 30), + serdes=JsonSerDes(), + ), + ) + context.step(submit(callback.callback_id), name="submit") + try: + result = callback.result() + except CallbackError: + result = outcome + context.wait(Duration.from_seconds(1), name="first-replay") + context.wait(Duration.from_seconds(1), name="second-replay") + assert isinstance(result, str) + return result + + store = FileSystemExecutionStore.create(tmp_path) if filesystem else None + with DurableFunctionTestRunner( + handler=durable_execution(handler, plugins=[observer]), + store=store, + skip_time=False, + poll_interval=0.01, + execution_timeout=25, + ) as runner: + arn = runner.run_async(input="{}") + try: + assert submitter_entered.wait(5) + callback_id = published_callbacks.get(timeout=5) + if outcome == "success": + runner.send_callback_success(callback_id, result=b'"target"') + elif outcome == "failure": + runner.send_callback_failure( + callback_id, error=ErrorObject.from_message("early failure") + ) + expected_event = { + "success": "CallbackSucceeded", + "failure": "CallbackFailed", + "timeout": "CallbackTimedOut", + }[outcome] + deadline = time.monotonic() + 5 + while time.monotonic() < deadline: + events = runner.get_execution_history( + arn, include_execution_data=True + ).events + if any(event.event_type == expected_event for event in events): + break + time.sleep(0.01) + else: + raise AssertionError("Callback did not complete during submission") + assert not any( + event.event_type == "InvocationCompleted" for event in events + ) + finally: + release_submitter.set() + result = runner.wait_for_result(arn, timeout=15) + + assert result.status.value == "SUCCEEDED" + assert result.result is not None + assert json.loads(result.result) == ("target" if outcome == "success" else outcome) + assert len(submission_calls) == 1 + assert len(observer.invocations) >= 3 + assert len(observer.ends) == 1 + target_id = observer.ends[0].operation_id + assert all( + target_id not in invocation.updated_operations + for invocation in observer.invocations[1:] + ) diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py index cd2fd8743..04334b6db 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py @@ -9,8 +9,10 @@ import pytest from aws_durable_execution_sdk_python.lambda_service import ( + CallbackDetails, ChainedInvokeOptions, ErrorObject, + Operation as ServiceOperation, OperationStatus, OperationType, StepDetails, @@ -849,6 +851,65 @@ def test_create_callback_failed(): assert event.callback_failed_details.error.payload.message == "Callback failed" +@pytest.mark.parametrize("include_data", [False, True]) +@pytest.mark.parametrize("has_error", [False, True]) +def test_callback_failure_history_projection_preserves_sdk_state( + include_data, has_error +): + error = ErrorObject.from_message("details") if has_error else None + operation = ServiceOperation( + operation_id="callback", + operation_type=OperationType.CALLBACK, + status=OperationStatus.FAILED, + callback_details=CallbackDetails(callback_id="callback-id", error=error), + ) + context = EventCreationContext.create( + operation=operation, + event_id=3, + durable_execution_arn="arn:test", + start_input=StartDurableExecutionInput( + account_id="123", + function_name="test", + function_qualifier="$LATEST", + execution_name="test", + execution_timeout_seconds=300, + execution_retention_period_days=7, + ), + include_execution_data=include_data, + ) + event = Event.create_callback_event(context) + expected = {"Truncated": not (include_data and not has_error)} + if has_error or include_data: + expected["Payload"] = error.to_dict() if error else {} + assert event.callback_failed_details.error.to_dict() == expected + assert operation.callback_details.error is error + + +def test_callback_timeout_history_projection_is_unchanged(): + operation = ServiceOperation( + operation_id="callback", + operation_type=OperationType.CALLBACK, + status=OperationStatus.TIMED_OUT, + callback_details=CallbackDetails(callback_id="callback-id", error=None), + ) + context = EventCreationContext.create( + operation=operation, + event_id=3, + durable_execution_arn="arn:test", + start_input=StartDurableExecutionInput( + account_id="123", + function_name="test", + function_qualifier="$LATEST", + execution_name="test", + execution_timeout_seconds=300, + execution_retention_period_days=7, + ), + include_execution_data=True, + ) + event = Event.create_callback_event(context) + assert event.callback_timed_out_details.error.to_dict() == {"Truncated": True} + + def test_create_callback_timed_out(): operation = create_mock_operation("callback-1", status=OperationStatus.TIMED_OUT) error_obj = ErrorObject.from_message("Callback timed out") diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py index bde5e3d4b..353d320d4 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py @@ -3,6 +3,7 @@ import json from dataclasses import replace from datetime import datetime, timezone +from threading import Event, Thread from unittest.mock import patch, Mock import pytest @@ -955,6 +956,38 @@ def test_from_dict_with_none_result(): # region callback +@pytest.mark.parametrize("outcome", ["success", "failure", "timeout"]) +def test_callback_completion_records_only_successful_state_changes(outcome): + """Preserve payloads, token versions and consumed metadata on rejection.""" + operation = Operation( + operation_id="callback", + operation_type=OperationType.CALLBACK, + status=OperationStatus.STARTED, + callback_details=CallbackDetails(callback_id="callback-id"), + ) + execution = Execution("test-arn", _make_start_input(), [operation]) + complete = getattr(execution, f"complete_callback_{outcome}") + payload = b'"result"' if outcome == "success" else ErrorObject.from_message("error") + token_version = execution.token_sequence + result = complete("callback-id", payload) + + assert execution.updated_operation_ids == ["callback"] + assert execution.token_sequence == token_version + assert execution.seq_counter == 1 + assert result.callback_details.result == ( + '"result"' if outcome == "success" else None + ) + assert result.callback_details.error == (None if outcome == "success" else payload) + restored = Execution.from_json_dict(execution.to_json_dict()) + assert restored.updated_operation_ids == ["callback"] + execution.mark_state_delivered() + with pytest.raises(IllegalStateException, match="not in STARTED state"): + complete("callback-id", payload) + assert execution.updated_operation_ids == [] + assert execution.seq_counter == 1 + assert execution.token_sequence == token_version + + def test_find_callback_operation_not_found(): """Test find_callback_operation raises exception when callback not found.""" execution = Execution("test-arn", Mock(), []) @@ -1690,6 +1723,61 @@ def test_record_invocation_completion_keeps_updated_operation_ids(): assert execution.updated_operation_ids == [] +@pytest.mark.parametrize("complete_before_advance", [False, True]) +def test_checkpoint_consumes_only_updates_covered_by_its_watermark( + complete_before_advance, +): + """A late update survives reads and retries of an older state delivery.""" + execution = Execution( + "test-arn", + _make_start_input(), + [ + Operation( + operation_id=name, + operation_type=OperationType.CALLBACK, + status=OperationStatus.STARTED, + callback_details=CallbackDetails(callback_id=name), + ) + for name in ["delivered", "later"] + ], + ) + execution.complete_callback_success("delivered", b"first") + response = OperationPaginatorState.pin(execution) + ready = Event() + finished = Event() + + def complete_later(): + assert ready.wait(5) + execution.complete_callback_failure("later", ErrorObject.from_message("later")) + finished.set() + + worker = Thread(target=complete_later) + worker.start() + if complete_before_advance: + ready.set() + assert finished.wait(5) + expected = ["delivered", "later"] if complete_before_advance else ["delivered"] + assert execution.updated_operation_ids == expected + response.page(None, max_size_bytes=1024 * 1024) + assert execution.updated_operation_ids == expected + + response.advance_handler_seen(1) + ready.set() + assert finished.wait(5) + worker.join(timeout=5) + assert not worker.is_alive() + assert execution.updated_operation_ids == ["later"] + assert execution.handler_seen_seq == 1 + assert execution.token_sequence == 0 + assert execution.seq_counter == 2 + # An idempotent/older delivery cannot consume the later completion. + response.advance_handler_seen(1) + response.advance_handler_seen(0) + assert execution.updated_operation_ids == ["later"] + OperationPaginatorState.pin(execution).advance_handler_seen(2) + assert execution.updated_operation_ids == [] + + def test_function_arn_is_qualified_with_the_executed_version(): execution = Execution.new(_make_start_input()) execution.region = "ap-southeast-2" diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/executor_checkpoint_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/executor_checkpoint_test.py index c789dcd74..fb21cf68f 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/executor_checkpoint_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/executor_checkpoint_test.py @@ -17,6 +17,7 @@ import pytest from aws_durable_execution_sdk_python.execution import InvocationStatus from aws_durable_execution_sdk_python.lambda_service import ( + CallbackDetails, ErrorObject, Operation as SvcOperation, OperationAction, @@ -126,6 +127,49 @@ def test_empty_poll_returns_empty_operations_and_advances_token(): assert CheckpointToken.from_str(response.checkpoint_token).token_sequence == 1 +@pytest.mark.parametrize("rejection", ["token", "operation"]) +def test_rejected_checkpoint_retains_undelivered_callback_update(rejection): + executor, store, execution, token = _make_executor_with_started_execution() + execution.operations.append( + SvcOperation( + operation_id="callback", + operation_type=OperationType.CALLBACK, + status=OperationStatus.STARTED, + callback_details=CallbackDetails(callback_id="callback-id"), + ) + ) + execution.complete_callback_success("callback-id", b"result") + store.save(execution) + before = ( + execution.token_sequence, + execution.handler_seen_seq, + execution.seq_counter, + ) + assert execution.updated_operation_ids == ["callback"] + with pytest.raises(InvalidParameterValueException): + executor.checkpoint_execution( + execution_arn=execution.durable_execution_arn, + checkpoint_token="invalid" if rejection == "token" else token, + updates=( + [ + OperationUpdate( + operation_id="callback", + operation_type=OperationType.CALLBACK, + action=OperationAction.SUCCEED, + ) + ] + if rejection == "operation" + else [] + ), + ) + assert execution.updated_operation_ids == ["callback"] + assert ( + execution.token_sequence, + execution.handler_seen_seq, + execution.seq_counter, + ) == before + + # endregion # region: Non-empty checkpoint returns only the delta diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py index 746f083d6..f3f8a9131 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py @@ -1461,6 +1461,36 @@ def test_event_error_with_payload_only(): } +@pytest.mark.parametrize( + "payload", + [ + {}, + {"ErrorMessage": ""}, + {"ErrorType": ""}, + {"ErrorData": ""}, + {"StackTrace": []}, + ], +) +def test_event_error_roundtrip_preserves_present_payload(payload): + wire = {"Payload": payload, "Truncated": False} + assert EventError.from_dict(wire).to_dict() == wire + + +@pytest.mark.parametrize("truncated", [False, True]) +@pytest.mark.parametrize("payload_field", [{}, {"Payload": None}]) +def test_event_error_absent_or_null_payload_retains_existing_meaning( + payload_field, truncated +): + parsed = EventError.from_dict({**payload_field, "Truncated": truncated}) + assert parsed.payload is None + assert parsed.to_dict() == {"Truncated": truncated} + + +def test_event_error_empty_payload_retains_truncation_flag(): + wire = {"Payload": {}, "Truncated": True} + assert EventError.from_dict(wire).to_dict() == wire + + # Tests for RetryDetails def test_retry_details_serialization(): """Test RetryDetails from_dict/to_dict round-trip.""" diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py index ab75f0d24..0662c07f1 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py @@ -1116,6 +1116,82 @@ def test_durable_child_context_test_runner_init_with_args( # Tests for DurableFunctionCloudTestRunner and from_execution_history +@pytest.mark.parametrize("decode_wire", [False, True]) +@pytest.mark.parametrize( + "payload", + [ + {}, + {"ErrorMessage": ""}, + {"ErrorType": ""}, + {"ErrorData": ""}, + {"StackTrace": []}, + {"ErrorMessage": "failure"}, + ], +) +def test_callback_history_result_preserves_wire_and_canonical_error( + payload, decode_wire +): + from aws_durable_execution_sdk_python.lambda_service import ErrorObject + from aws_durable_execution_sdk_python_testing.model import ( + CallbackFailedDetails, + CallbackStartedDetails, + Event, + EventError, + GetDurableExecutionResponse, + ) + + timestamp = datetime.datetime(2026, 10, 7, tzinfo=datetime.UTC) + caller_error = ErrorObject.from_message("Callback failed") + execution = GetDurableExecutionResponse( + durable_execution_arn="arn:execution", + durable_execution_name="execution", + function_arn="arn:function", + status="FAILED", + start_timestamp=timestamp, + error=caller_error, + ) + history = GetDurableExecutionHistoryResponse( + events=[ + Event( + event_type="CallbackStarted", + event_timestamp=timestamp, + event_id=1, + operation_id="callback", + name="callback", + callback_started_details=CallbackStartedDetails( + callback_id="callback-id" + ), + ), + Event( + event_type="CallbackFailed", + event_timestamp=timestamp, + event_id=2, + operation_id="callback", + name="callback", + callback_failed_details=CallbackFailedDetails( + error=EventError( + payload=ErrorObject.from_dict(payload), truncated=False + ) + ), + ), + ] + ) + if decode_wire: + history = GetDurableExecutionHistoryResponse.from_dict(history.to_dict()) + assert history.events[-1].callback_failed_details.error.to_dict() == { + "Payload": payload, + "Truncated": False, + } + result = DurableFunctionTestResult.from_execution_history(execution, history) + callback = result.get_callback("callback") + assert callback.status is OperationStatus.FAILED + assert (callback.error.to_dict() if callback.error is not None else None) == ( + payload if payload else None + ) + assert result.status is InvocationStatus.FAILED + assert result.error is caller_error + + def test_durable_function_test_result_from_execution_history(): """Test DurableFunctionTestResult.from_execution_history factory method.""" import datetime diff --git a/packages/aws-durable-execution-sdk-python/README.md b/packages/aws-durable-execution-sdk-python/README.md index 18bc0d5e7..4b2463e52 100644 --- a/packages/aws-durable-execution-sdk-python/README.md +++ b/packages/aws-durable-execution-sdk-python/README.md @@ -85,9 +85,8 @@ constraint applies to the combined `plugins=[...]` argument and Core 2.1+ with OTel 1.1+ rejects both at cold start with `PluginLoadError` naming the conflicting views; keep only one. Choose Invocation for work within each Lambda invocation or Execution for logical operations across the durable execution. -Unrelated instrumentation plugins can run alongside either view. Existing valid -registrations remain supported with OTel 1.1 on older core 2.0.x; those cores do -not implement the new exclusivity validation. +Unrelated instrumentation plugins can run alongside either view. OTel 1.1 requires core 2.1 or later for the invocation-worker lifecycle as well +as view-exclusivity validation. Release core 2.1 before releasing OTel 1.1. Plugin authors explicitly opt in by declaring `__durable_registration_api__ = 1` on a plugin class. That class and its subclasses @@ -112,6 +111,40 @@ in their hierarchy retain their existing attributes/helpers. The generic plugin neither attribute nor hook, and provider API version 1 and existing plugin lifecycle order are unchanged. +### Invocation worker lifecycle + +The existing `on_invocation_start` and `on_invocation_end` hooks run on the +handler's invocation worker, both in registration order. Start hooks finish +before background checkpoint processing begins. The same worker runs the handler +(including its `finally` blocks), prepares the existing serialized output or +error, checkpoints large results when necessary, joins registered branches while +checkpointing is available, stops and waits for checkpoint processing, then calls +End before returning the outcome. The caller shuts down the handler executor. +There is no separate handler-context hook or context-manager plugin API. + +When plugins are registered, the worker begins with a copy of the caller's +context-variable bindings. Successful Start bindings flow into later Start hooks, +the handler, serialization and resource cleanup. A later successful Start may +replace an earlier binding. If a Start hook raises, its new bindings are discarded +for subsequent work; its End still runs in the original Context so its tokens can +be reset. End hooks retain forward registration order, not reverse stack order, +so they must not rely on observing a stack-like unwind of other plugins' contexts. +Each SDK-managed `map` or `parallel` branch admission, including an in-process +resume, receives a fresh copy of the coordinator's bindings when plugins are +registered. Branch changes remain local even when pool threads are reused. +User-created threads retain normal Python context-variable behavior; without +registered plugins, SDK branch submission keeps its existing behavior. +The worker's invocation context is discarded on return, including after plugin +cleanup failures, leaving the host's bindings unchanged. This isolates bindings, +not mutations to shared objects or external side effects. Without plugins, the +handler retains its existing fresh-worker context behavior. + +These lifecycle guarantees require core 2.1.0 or later. The OTel 1.1 plugin +requires that core version and uses the existing Start/End hooks. Upgrading the +core alone with OTel 1.0 isolates worker bindings, but does not add the newer +Invocation-view fallback to that older plugin; upgrade both packages for it. +Provider API version 1 and the independent registration API are unchanged. + ## 🚀 Quick Start Install the execution SDK: diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/concurrency/executor.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/concurrency/executor.py index 2a39c4eff..bafa43b19 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/concurrency/executor.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/concurrency/executor.py @@ -2,6 +2,7 @@ from __future__ import annotations +import contextvars import heapq import logging import queue @@ -291,13 +292,25 @@ def execute( def submit(branch: Branch[CallableType, ResultType]) -> None: branch.start() - pool.submit( - self._branch_worker, - execution_state, - executor_context, - events, - branch.executable, - ) + if execution_state._plugin_executor._plugins: # noqa: SLF001 + # Every admission/resume gets its own Context. Branch bindings + # cannot leak into siblings, the coordinator, or reused workers. + pool.submit( + contextvars.copy_context().run, + self._branch_worker, + execution_state, + executor_context, + events, + branch.executable, + ) + else: + pool.submit( + self._branch_worker, + execution_state, + executor_context, + events, + branch.executable, + ) try: # Only rebuild the items snapshot after a terminal event changes diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/execution.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/execution.py index 8ca5ef903..e59820007 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/execution.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/execution.py @@ -1,10 +1,11 @@ from __future__ import annotations import contextlib +import contextvars import functools import json import logging -from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import Future, ThreadPoolExecutor, wait from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any @@ -188,9 +189,9 @@ def durable_execution( logger.debug("Starting durable execution handler...") - plugin_executor = PluginExecutor(load_configured_plugins(plugins)) + configured_plugins = load_configured_plugins(plugins) + plugin_executor = PluginExecutor(configured_plugins) - @plugin_executor.handle_durable_output def wrapper(event: Any, context: LambdaContext) -> MutableMapping[str, Any]: invocation_input: DurableExecutionInvocationInput service_client: DurableServiceClient @@ -279,50 +280,13 @@ def wrapper(event: Any, context: LambdaContext) -> MutableMapping[str, Any]: ), ) - # Use ThreadPoolExecutor for concurrent execution of user code and background checkpoint processing - with ( - ThreadPoolExecutor( - max_workers=2, thread_name_prefix="dex-handler" - ) as executor, - contextlib.closing(execution_state) as execution_state, - ): - execution_operation = execution_state.get_execution_operation() - - # execute the plugins - plugin_executor.on_invocation_start( - execution_arn=invocation_input.durable_execution_arn, - lambda_context=context, - execution_start_time=( - execution_operation.start_timestamp - if execution_operation is not None - else None - ), - is_first_invocation=not has_prior_operations, - execution_input=input_event, - # Read the map through a callable rather than snapshotting it - # here: the invocation-end hook needs the state as of the end of - # the invocation, and neither hook pays for the conversion until - # a plugin actually reads it. - operations_provider=lambda: execution_state.operations, - updated_operation_ids=invocation_input.updated_operation_ids, - ) - # Thread 1: Run background checkpoint processing - executor.submit(execution_state.checkpoint_batches_forever) - - # Thread 2: Execute user function + def invoke_handler() -> MutableMapping[str, Any]: logger.debug( "%s entering user-space...", invocation_input.durable_execution_arn ) - user_future = executor.submit(func, input_event, durable_context) - - logger.debug( - "%s waiting for user code completion...", - invocation_input.durable_execution_arn, - ) - try: # Background checkpointing errors will propagate through CompletionEvent.wait() as BackgroundThreadError - result = user_future.result() + result = func(input_event, durable_context) # done with userland logger.debug( @@ -458,6 +422,78 @@ def wrapper(event: Any, context: LambdaContext) -> MutableMapping[str, Any]: return result + def invoke_worker() -> MutableMapping[str, Any]: + # This task owns Start, user cleanup, output preparation, resource + # cleanup and End. The caller owns the executor that runs this task. + with plugin_executor.run(): + checkpoint_future: Future[None] | None = None + + def wait_for_checkpoint() -> None: + if checkpoint_future is not None: + # Match the former executor join: observed checkpoint + # errors retain their existing classification paths. + wait((checkpoint_future,)) + + try: + with contextlib.ExitStack() as resources: + # LIFO cleanup: branches join with checkpointing alive; + # then close stops checkpointing and the wait completes. + resources.callback(wait_for_checkpoint) + resources.callback( + plugin_executor._run_in_invocation_context, + execution_state.close, + ) + execution_operation = execution_state.get_execution_operation() + + # execute the plugins + plugin_executor.on_invocation_start( + execution_arn=invocation_input.durable_execution_arn, + lambda_context=context, + execution_start_time=( + execution_operation.start_timestamp + if execution_operation is not None + else None + ), + is_first_invocation=not has_prior_operations, + execution_input=input_event, + # Read the map through a callable rather than snapshotting it + # here: the invocation-end hook needs the state as of the end of + # the invocation, and neither hook pays for the conversion until + # a plugin actually reads it. + operations_provider=lambda: execution_state.operations, + updated_operation_ids=invocation_input.updated_operation_ids, + ) + # No checkpoint work starts until all Start hooks finish. + checkpoint_future = executor.submit( + execution_state.checkpoint_batches_forever + ) + output = plugin_executor._run_in_invocation_context( + invoke_handler + ) + plugin_executor.on_invocation_end( + DurableExecutionInvocationOutput.from_dict(output) + ) + return output + except Exception as error: + plugin_executor.on_invocation_end( + DurableExecutionInvocationOutput.create_retry( + ErrorObject.from_exception(error) + ) + ) + raise + + with ThreadPoolExecutor( + max_workers=2, thread_name_prefix="dex-handler" + ) as executor: + if configured_plugins: + invocation_future = executor.submit( + contextvars.copy_context().run, invoke_worker + ) + else: + # Preserve the original fresh-worker context without plugins. + invocation_future = executor.submit(invoke_worker) + return invocation_future.result() + return wrapper diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py index e9549d308..ed3792102 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py @@ -1,15 +1,16 @@ from __future__ import annotations import contextlib +import contextvars import copy import datetime -import functools import logging from collections.abc import Mapping, Sequence from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from enum import Enum -from typing import Any, Callable, MutableMapping, cast +from threading import Lock +from typing import Any, Callable, cast from aws_durable_execution_sdk_python.identifier import OperationIdentifier from aws_durable_execution_sdk_python.lambda_service import ( @@ -466,9 +467,17 @@ def __init__(self, plugins: list[DurableInstrumentationPlugin] | None): self._executor: ThreadPoolExecutor | None = None self._invocation_status: InvocationStartInfo | None = None self._operations_provider: Callable[[], Mapping[str, Operation]] | None = None + # Non-None only after a Start hook fails: later setup and invocation + # work use its pre-hook snapshot. End still uses each hook's token owner. + self._startup_context: contextvars.Context | None = None + self._invocation_contexts: list[contextvars.Context | None] = [] + self._reported_terminal_updates: set[tuple[str, OperationStatus]] = set() + self._terminal_updates_lock = Lock() @contextlib.contextmanager def run(self): + with self._terminal_updates_lock: + self._reported_terminal_updates.clear() if self._plugins: self._executor = ThreadPoolExecutor( max_workers=1, @@ -479,12 +488,16 @@ def run(self): finally: self._invocation_status = None self._operations_provider = None + self._startup_context = None + self._invocation_contexts.clear() # Shut down the thread pool, waiting for pending tasks to complete. if self._executor: self._executor.shutdown(wait=True) + with self._terminal_updates_lock: + self._reported_terminal_updates.clear() @staticmethod - def _dispatch_plugin(plugin: DurableInstrumentationPlugin, info) -> None: + def _dispatch_plugin(plugin: DurableInstrumentationPlugin, info) -> bool: """Invoke the appropriate plugin callback. Runs inside the thread pool.""" try: match info: @@ -507,18 +520,54 @@ def _dispatch_plugin(plugin: DurableInstrumentationPlugin, info) -> None: except Exception: # log and ignore the exception logger.exception("Plugin %s exception ignored", plugin.__class__.__name__) + return False + return True def execute_plugins(self, info, sync): if not self._executor: return - for plugin in self._plugins: - if sync: - # this is called synchronously, so plugins will be able to manipulate thread local objects - self._dispatch_plugin(plugin, info) + if sync and isinstance(info, InvocationStartInfo): + self._startup_context = None + self._invocation_contexts.clear() + for index, plugin in enumerate(self._plugins): + if sync and isinstance(info, InvocationStartInfo): + owner = self._startup_context + before = ( + owner.copy() if owner is not None else contextvars.copy_context() + ) + self._invocation_contexts.append(owner) + succeeded = ( + owner.run(self._dispatch_plugin, plugin, info) + if owner is not None + else self._dispatch_plugin(plugin, info) + ) + if not succeeded: + # A failing hook may have left new bindings with no reset token. + # Continue setup and the handler in the pre-hook snapshot. + self._startup_context = before + elif sync: + # End hooks must reset tokens in the Context that created them, + # even when a failed start hook moved later setup to a snapshot. + owner = ( + self._invocation_contexts[index] + if isinstance(info, InvocationEndInfo) + and index < len(self._invocation_contexts) + else None + ) + if owner is not None: + owner.run(self._dispatch_plugin, plugin, info) + else: + self._dispatch_plugin(plugin, info) else: # this is called asynchronously, so plugins cannot manipulate thread local objects self._executor.submit(self._dispatch_plugin, plugin, info) + def _run_in_invocation_context(self, invoke: Callable[[], Any]) -> Any: + """Continue in the pre-hook Context only after a failed Start hook.""" + if self._startup_context is not None: + return self._startup_context.run(invoke) + return invoke() + def _snapshot_operation_infos( self, operations_provider: Callable[[], Mapping[str, Operation]] | None, @@ -773,6 +822,15 @@ def on_operation_update( ) for operation in updated_operations: if self._is_terminal_status(operation.status): + # Replay delivery and checkpoint responses can report the same + # completion. Deduplicate actual notifications, not state that + # may have arrived without an UpdatedOperationIds notification. + if self._plugins: + key = (operation.operation_id, operation.status) + with self._terminal_updates_lock: + if key in self._reported_terminal_updates: + continue + self._reported_terminal_updates.add(key) self.execute_plugins( OperationEndInfo( operation_id=operation.operation_id, @@ -836,28 +894,3 @@ def _is_terminal_status(status): OperationStatus.CANCELLED, OperationStatus.STOPPED, ] - - @property - def handle_durable_output(self): - def decorator(func: Callable[[Any, LambdaContext], MutableMapping[str, Any]]): - @functools.wraps(func) - def wrapper(event: Any, context: LambdaContext): - with self.run(): - try: - output = func(event, context) - - self.on_invocation_end( - output=DurableExecutionInvocationOutput.from_dict(output), - ) - return output - except Exception as e: - self.on_invocation_end( - output=DurableExecutionInvocationOutput.create_retry( - ErrorObject.from_exception(e) - ), - ) - raise - - return wrapper - - return decorator diff --git a/packages/aws-durable-execution-sdk-python/tests/concurrency_test.py b/packages/aws-durable-execution-sdk-python/tests/concurrency_test.py index a93135663..c30980c73 100644 --- a/packages/aws-durable-execution-sdk-python/tests/concurrency_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/concurrency_test.py @@ -1,5 +1,6 @@ """Tests for the concurrency module.""" +import contextvars import hashlib import json import queue @@ -77,7 +78,10 @@ from aws_durable_execution_sdk_python.operation.parallel import ( ParallelExecutor, ) -from aws_durable_execution_sdk_python.plugin import PluginExecutor +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + PluginExecutor, +) from aws_durable_execution_sdk_python.state import ( CheckpointedResult, ExecutionState, @@ -5275,3 +5279,44 @@ def predicate(s: CompletionStatus) -> CompletionDecision: # endregion Custom completion predicate (should_complete) integration tests + + +def test_instrumented_executor_isolates_bindings_on_a_reused_worker() -> None: + """Exercise actual admission/worker execution independently of durable I/O.""" + binding = contextvars.ContextVar("executor-unit-binding", default="empty") + seen: list[tuple[str, int]] = [] + + class RecordingExecutor(ConcurrentExecutor[Callable[[], str], str]): + def _execute_item_in_child_context( + self, + executor_context: DurableContext, + executable: Executable[Callable[[], str]], + ) -> str: + seen.append((binding.get(), threading.get_ident())) + binding.set(f"branch-{executable.index}") + return executable.func() + + executor = RecordingExecutor( + executables=[Executable(index, lambda: "ok") for index in range(2)], + max_concurrency=1, + completion_config=CompletionConfig(min_successful=2), + sub_type_top=OperationSubType.PARALLEL, + sub_type_iteration=OperationSubType.PARALLEL_BRANCH, + name_prefix="branch-", + serdes=None, + operation_id_namespace=_StubNamespace(), + ) + state = Mock(spec=ExecutionState) + state._plugin_executor = PluginExecutor([DurableInstrumentationPlugin()]) + token = binding.set("coordinator") + try: + result = executor.execute(state, Mock(spec=DurableContext)) + assert result.get_results() == ["ok", "ok"] + assert binding.get() == "coordinator" + assert [value for value, _thread in seen] == ["coordinator", "coordinator"] + assert seen[0][1] == seen[1][1] != threading.get_ident() + finally: + for call in state.register_branch_pool.call_args_list: + call.args[0].shutdown(wait=True) + binding.reset(token) + assert binding.get() == "empty" diff --git a/packages/aws-durable-execution-sdk-python/tests/e2e/branch_worker_context_test.py b/packages/aws-durable-execution-sdk-python/tests/e2e/branch_worker_context_test.py new file mode 100644 index 000000000..a9104a745 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/e2e/branch_worker_context_test.py @@ -0,0 +1,129 @@ +"""Registered Start bindings reach independently resumed public branch workers.""" + +from __future__ import annotations + +import contextvars +import json +import threading +from typing import Any + +import pytest +from aws_durable_execution_sdk_python import DurableContext, durable_execution +from aws_durable_execution_sdk_python.config import Duration, MapConfig, ParallelConfig +from aws_durable_execution_sdk_python.concurrency.models import BatchResult +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationEndInfo, + InvocationStartInfo, + InvocationStatus, +) +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner + + +@pytest.mark.parametrize("kind", ["map", "parallel"]) +@pytest.mark.parametrize("concurrency", [1, 2]) +@pytest.mark.parametrize("failed_start", [False, True]) +def test_partial_branch_resume_uses_fresh_successful_start_bindings( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + kind: str, + concurrency: int, + failed_start: bool, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + marker = contextvars.ContextVar("resumed-branch", default="host") + partial = contextvars.ContextVar[str]("partial-start") + observations: list[tuple[int, str]] = [] + visits = {0: 0, 1: 0} + bodies: list[int] = [] + statuses: list[InvocationStatus] = [] + lock = threading.Lock() + barrier = threading.Barrier(2) if concurrency == 2 else None + + class Bind(DurableInstrumentationPlugin): + token: contextvars.Token[str] | None = None + + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + self.token = marker.set("plugin") + + def on_invocation_end(self, info: InvocationEndInfo) -> None: + statuses.append(info.status) + assert self.token is not None + marker.reset(self.token) + self.token = None + + class Broken(DurableInstrumentationPlugin): + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + marker.set("failed") + partial.set("failed") + raise ValueError("expected failed Start") + + plugin = Bind() + plugins: list[DurableInstrumentationPlugin] = [plugin] + if failed_start: + plugins.append(Broken()) + + def branch(child: DurableContext, index: int) -> str: + with lock: + visits[index] += 1 + first_visit = visits[index] == 1 + observations.append((index, marker.get())) + with pytest.raises(LookupError): + partial.get() + marker.set(f"branch-{index}") + if barrier is not None and first_visit: + barrier.wait(timeout=10) + + def step(_step: Any) -> str: + with lock: + bodies.append(index) + return f"saved-{index}" + + saved = child.step(step, name="save") + child.wait(Duration.from_seconds(index + 1), name="branch-wait") + return saved + + def handler(_event: Any, durable: DurableContext) -> list[str]: + assert marker.get() == "plugin" + result: BatchResult[str] + if kind == "map": + result = durable.map( + [0, 1], + lambda child, item, index, items: branch(child, index), + name="mapped", + config=MapConfig(max_concurrency=concurrency), + ) + else: + result = durable.parallel( + [lambda child: branch(child, 0), lambda child: branch(child, 1)], + name="parallel", + config=ParallelConfig(max_concurrency=concurrency), + ) + assert marker.get() == "plugin" + durable.wait(Duration.from_seconds(1), name="after-batch") + return result.get_results() + + wrapped = durable_execution(handler, plugins=plugins) + with DurableFunctionTestRunner(handler=wrapped) as runner: + result = runner.run(input="{}", timeout=30) + assert result.status.value == "SUCCEEDED" + assert json.loads(result.result) == ["saved-0", "saved-1"] + assert sorted(bodies) == [0, 1] + assert all(count >= 2 for count in visits.values()) + assert all(binding == "plugin" for _, binding in observations) + assert InvocationStatus.PENDING in statuses + assert statuses[-1] is InvocationStatus.SUCCEEDED + assert marker.get() == "host" and plugin.token is None + with pytest.raises(LookupError): + partial.get() + errors = [ + record + for record in caplog.records + if record.exc_info and record.name == "aws_durable_execution_sdk_python.plugin" + ] + assert len(errors) == (len(statuses) if failed_start else 0) + assert all( + record.exc_info is not None + and str(record.exc_info[1]) == "expected failed Start" + for record in errors + ) diff --git a/packages/aws-durable-execution-sdk-python/tests/e2e/handler_invocation_context_int_test.py b/packages/aws-durable-execution-sdk-python/tests/e2e/handler_invocation_context_int_test.py new file mode 100644 index 000000000..59f5c9b52 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/e2e/handler_invocation_context_int_test.py @@ -0,0 +1,79 @@ +"""Invocation context isolation across real suspension and replay.""" + +from __future__ import annotations + +import contextvars +import json +from typing import Any + +import pytest +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner + +from aws_durable_execution_sdk_python import DurableContext, durable_execution +from aws_durable_execution_sdk_python.config import Duration +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationEndInfo, + InvocationStartInfo, +) + + +@pytest.mark.parametrize("fail_after_resume", [False, True]) +@pytest.mark.parametrize("reverse", [False, True]) +def test_invocation_plugins_restore_host_context_across_resume( + monkeypatch: pytest.MonkeyPatch, fail_after_resume: bool, reverse: bool +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + marker = contextvars.ContextVar("invocation-host", default="host") + boundaries: list[tuple[str, str]] = [] + body_calls: list[str] = [] + + class ScopePlugin(DurableInstrumentationPlugin): + def __init__(self, name: str) -> None: + self.name = name + self.token: contextvars.Token[str] | None = None + + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + self.token = marker.set(self.name) + + def on_invocation_end(self, _info: InvocationEndInfo) -> None: + assert self.token is not None + marker.reset(self.token) + self.token = None + + names = ["first", "second"] + if reverse: + names.reverse() + + def body(_event: Any, context: DurableContext) -> str: + assert marker.get() == names[-1] + + def step(_step_context: Any) -> str: + body_calls.append("step") + return "saved" + + saved = context.step(step, name="before-wait") + context.wait(Duration.from_seconds(1), name="resume") + assert marker.get() == names[-1] + if fail_after_resume: + raise ValueError("failure after resume") + return saved + + durable_handler = durable_execution( + body, plugins=[ScopePlugin(name) for name in names] + ) + + def host(event: Any, context: Any) -> Any: + before = marker.get() + try: + return durable_handler(event, context) + finally: + boundaries.append((before, marker.get())) + + with DurableFunctionTestRunner(handler=host) as runner: + result = runner.run(input="{}", timeout=15) + assert result.status.value == ("FAILED" if fail_after_resume else "SUCCEEDED") + if not fail_after_resume: + assert json.loads(result.result) == "saved" + assert body_calls == ["step"] + assert boundaries == [("host", "host"), ("host", "host")] diff --git a/packages/aws-durable-execution-sdk-python/tests/handler_worker_context_test.py b/packages/aws-durable-execution-sdk-python/tests/handler_worker_context_test.py new file mode 100644 index 000000000..8cb081029 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/handler_worker_context_test.py @@ -0,0 +1,589 @@ +"""Real public invocation-worker lifecycle and ContextVar ownership controls.""" + +from __future__ import annotations + +import contextvars +import threading +from concurrent.futures import ThreadPoolExecutor +from typing import Any +from unittest.mock import Mock + +import pytest + +from aws_durable_execution_sdk_python.context import DurableContext +from aws_durable_execution_sdk_python.exceptions import ( + CheckpointError, + CheckpointErrorCategory, + InvocationError, + SuspendExecution, +) +from aws_durable_execution_sdk_python.execution import ( + DurableExecutionInvocationInputWithClient, + InitialExecutionState, + durable_execution, +) +from aws_durable_execution_sdk_python.lambda_service import ( + CheckpointOutput, + CheckpointUpdatedExecutionState, + DurableServiceClient, + ExecutionDetails, + Operation, + OperationStatus, + OperationType, + OperationUpdate, +) +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationEndInfo, + InvocationStartInfo, + InvocationStatus, +) +from aws_durable_execution_sdk_python.state import ExecutionState + + +def invocation(client: Any = None) -> tuple[Any, Any]: + client = client or Mock(spec=DurableServiceClient) + client.checkpoint.return_value = CheckpointOutput( + checkpoint_token="next", new_execution_state=CheckpointUpdatedExecutionState() + ) + event = DurableExecutionInvocationInputWithClient( + durable_execution_arn="test-arn/worker-lifecycle", + checkpoint_token="initial", + initial_execution_state=InitialExecutionState( + operations=[ + Operation( + operation_id="execution", + operation_type=OperationType.EXECUTION, + status=OperationStatus.STARTED, + execution_details=ExecutionDetails(input_payload="{}"), + ) + ], + next_marker="", + ), + service_client=client, + ) + context = Mock() + context.aws_request_id = "worker-request" + context.client_context = context.identity = context.invoked_function_arn = None + context._epoch_deadline_time_in_ms = 0 + context.tenant_id = None + return event, context + + +@pytest.mark.parametrize("outcome", ["success", "failure", "retry", "pending"]) +@pytest.mark.parametrize("mode", ["none", "healthy", "failed-start", "failed-end"]) +def test_public_worker_lifecycle_owns_context_and_outcome( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + outcome: str, + mode: str, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + marker = contextvars.ContextVar("worker-owner", default="fresh-worker") + phases: list[tuple[str, int, str]] = [] + boundary: list[tuple[str, int, str, str]] = [] + retry = InvocationError("original retry") + fail = ValueError("original body failure") + statuses: list[InvocationStatus] = [] + caller_thread = threading.get_ident() + + class Observer(ThreadPoolExecutor): + def submit(self, fn: Any, /, *args: Any, **kwargs: Any) -> Any: + role = ( + "checkpoint" + if getattr(fn, "__name__", "") == "checkpoint_batches_forever" + else "invocation" + ) + + def observed() -> Any: + before = marker.get() + try: + return fn(*args, **kwargs) + finally: + boundary.append((role, threading.get_ident(), before, marker.get())) + + return super().submit(observed) + + monkeypatch.setattr( + "aws_durable_execution_sdk_python.execution.ThreadPoolExecutor", Observer + ) + + class Plugin(DurableInstrumentationPlugin): + token: Any = None + + def on_invocation_start(self, info: InvocationStartInfo) -> None: + phases.append(("start", threading.get_ident(), marker.get())) + self.token = marker.set("start-binding") + if mode == "failed-start": + raise ValueError("startup failure") + + def on_invocation_end(self, info: InvocationEndInfo) -> None: + phases.append(("end", threading.get_ident(), marker.get())) + statuses.append(info.status) + assert self.token is not None + marker.reset(self.token) + self.token = None + if mode == "failed-end": + raise ValueError("end failure") + + def body(_event: Any, _context: DurableContext) -> str: + phases.append(("body", threading.get_ident(), marker.get())) + marker.set("body-binding") + try: + if outcome == "failure": + raise fail + if outcome == "retry": + raise retry + if outcome == "pending": + raise SuspendExecution("test suspension") + return "ok" + finally: + phases.append(("finally", threading.get_ident(), marker.get())) + + handler = durable_execution(body, plugins=[] if mode == "none" else [Plugin()]) + event, context = invocation() + token = marker.set("host") + try: + for _ in range(2): + if outcome == "retry": + with pytest.raises(InvocationError) as caught: + handler(event, context) + assert caught.value is retry + else: + result = handler(event, context) + assert ( + result["Status"] + == { + "success": "SUCCEEDED", + "failure": "FAILED", + "pending": "PENDING", + }[outcome] + ) + if outcome == "failure": + assert result["Error"]["ErrorMessage"] == "original body failure" + assert marker.get() == "host" + finally: + marker.reset(token) + per_call = 2 if mode == "none" else 4 + for offset in range(0, len(phases), per_call): + batch = phases[offset : offset + per_call] + assert [p[0] for p in batch] == ( + ["body", "finally"] + if mode == "none" + else ["start", "body", "finally", "end"] + ) + assert len({p[1] for p in batch}) == 1 + assert batch[0][1] != caller_thread + bodies = [p for p in phases if p[0] == "body"] + expected = ( + "fresh-worker" + if mode == "none" + else "host" + if mode == "failed-start" + else "start-binding" + ) + assert [p[2] for p in bodies] == [expected, expected] + for role, tid, before, after in boundary: + if role == "invocation": + assert after == ("body-binding" if mode == "none" else before) + if mode != "none": + assert ( + statuses + == [ + { + "success": InvocationStatus.SUCCEEDED, + "failure": InvocationStatus.FAILED, + "retry": InvocationStatus.RETRY, + "pending": InvocationStatus.PENDING, + }[outcome] + ] + * 2 + ) + assert "different Context" not in caplog.text + + +@pytest.mark.parametrize("bad_first", [False, True]) +@pytest.mark.parametrize("outcome", ["success", "failure", "retry", "pending"]) +def test_failed_start_preserves_clean_bindings_and_original_token_owners( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + bad_first: bool, + outcome: str, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + marker = contextvars.ContextVar("failed-start-owner", default="host") + added = contextvars.ContextVar[str]("failed-start-added") + events: list[tuple[str, str, int]] = [] + cleanup_contexts: list[str] = [] + original_close = ExecutionState.close + + def close(state: ExecutionState) -> None: + cleanup_contexts.append(marker.get()) + with pytest.raises(LookupError): + added.get() + original_close(state) + + monkeypatch.setattr(ExecutionState, "close", close) + + class Plugin(DurableInstrumentationPlugin): + def __init__(self, name: str) -> None: + self.name = name + self.token: contextvars.Token[str] | None = None + self.added: contextvars.Token[str] | None = None + + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + events.append(("start", self.name, threading.get_ident())) + self.token = marker.set(self.name) + if self.name == "bad": + self.added = added.set("partial") + raise ValueError("partial startup") + + def on_invocation_end(self, _info: InvocationEndInfo) -> None: + events.append(("end", self.name, threading.get_ident())) + assert self.token is not None + marker.reset(self.token) + if self.added is not None: + added.reset(self.added) + + failure = InvocationError("retry") + + def body(_event: Any, _context: DurableContext) -> str: + events.append(("body", "handler", threading.get_ident())) + assert marker.get() == "healthy" + with pytest.raises(LookupError): + added.get() + if outcome == "retry": + raise failure + if outcome == "failure": + raise ValueError("body") + if outcome == "pending": + raise SuspendExecution("test suspension") + return "ok" + + names = ["bad", "healthy"] if bad_first else ["healthy", "bad"] + handler = durable_execution(body, plugins=[Plugin(name) for name in names]) + event, context = invocation() + for _ in range(2): + if outcome == "retry": + with pytest.raises(InvocationError) as caught: + handler(event, context) + assert caught.value is failure + else: + assert ( + handler(event, context)["Status"] + == {"success": "SUCCEEDED", "failure": "FAILED", "pending": "PENDING"}[ + outcome + ] + ) + assert marker.get() == "host" + with pytest.raises(LookupError): + added.get() + assert cleanup_contexts == ["healthy"] * 2 + for offset in (0, 5): + batch = events[offset : offset + 5] + assert [(x[0], x[1]) for x in batch] == [("start", n) for n in names] + [ + ("body", "handler") + ] + [("end", n) for n in names] + assert len({x[2] for x in batch}) == 1 + errors = [ + r.exc_info + for r in caplog.records + if r.exc_info and r.name == "aws_durable_execution_sdk_python.plugin" + ] + assert len(errors) == 2 and all(str(e[1]) == "partial startup" for e in errors) + + +@pytest.mark.parametrize("result_kind", ["normal", "bad-json", "large", "large-error"]) +def test_worker_prepares_output_joins_branches_waits_checkpoint_then_ends( + monkeypatch: pytest.MonkeyPatch, + result_kind: str, +) -> None: + import aws_durable_execution_sdk_python.execution as execution + + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + events: list[tuple[str, int]] = [] + started = threading.Event() + checkpoint_started = threading.Event() + closing = threading.Event() + branch_done = threading.Event() + real_checkpoint = ExecutionState.checkpoint_batches_forever + real_close = ExecutionState.close + real_stop = ExecutionState.stop_checkpointing + real_dumps = execution.json.dumps + result: Any = { + "normal": {"ok": True}, + "bad-json": object(), + "large": "x" * 256, + "large-error": None, + }[result_kind] + + def record(name: str) -> None: + events.append((name, threading.get_ident())) + + def checkpoint(state: ExecutionState) -> None: + assert started.is_set() + record("checkpoint-start") + checkpoint_started.set() + try: + real_checkpoint(state) + finally: + record("checkpoint-end") + + def stop(state: ExecutionState) -> None: + assert branch_done.is_set() + record("checkpoint-stop") + real_stop(state) + + def close(state: ExecutionState) -> None: + record("close") + closing.set() + real_close(state) + + def dumps(value: Any, *args: Any, **kwargs: Any) -> str: + if value is result or ( + isinstance(value, dict) and value.get("Status") == "FAILED" + ): + record("serialize") + return real_dumps(value, *args, **kwargs) + + monkeypatch.setattr(ExecutionState, "checkpoint_batches_forever", checkpoint) + monkeypatch.setattr(ExecutionState, "stop_checkpointing", stop) + monkeypatch.setattr(ExecutionState, "close", close) + monkeypatch.setattr(execution.json, "dumps", dumps) + if result_kind in ("large", "large-error"): + monkeypatch.setattr(execution, "LAMBDA_RESPONSE_SIZE_LIMIT", 64) + + class Plugin(DurableInstrumentationPlugin): + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + record("start") + assert not checkpoint_started.is_set() + started.set() + + def on_invocation_end(self, info: InvocationEndInfo) -> None: + record("end") + assert branch_done.is_set() + assert any(x[0] == "checkpoint-end" for x in events) + assert info.status is ( + InvocationStatus.FAILED + if result_kind in ("bad-json", "large-error") + else InvocationStatus.SUCCEEDED + ) + + def body(_event: Any, ctx: DurableContext) -> Any: + assert checkpoint_started.wait(5) + record("body") + pool = ThreadPoolExecutor(max_workers=1) + ctx.state.register_branch_pool(pool) + + def late_branch() -> None: + assert closing.wait(5) + assert not ctx.state._checkpointing_stopped.is_set() + ctx.state.create_checkpoint( + OperationUpdate.create_execution_succeed(payload='"branch"'), + is_sync=True, + ) + record("branch-done") + branch_done.set() + + pool.submit(late_branch) + try: + if result_kind == "large-error": + raise ValueError("x" * 256) + return result + finally: + record("finally") + + handler = durable_execution(body, plugins=[Plugin()]) + event, context = invocation() + output = handler(event, context) + record("caller") + names = [x[0] for x in events] + assert ( + names.index("start") + < names.index("checkpoint-start") + < names.index("body") + < names.index("finally") + < names.index("serialize") + < names.index("close") + ) + assert ( + names.index("branch-done") + < names.index("checkpoint-stop") + < names.index("checkpoint-end") + < names.index("end") + < names.index("caller") + ) + worker_events = { + tid + for name, tid in events + if name + in {"start", "body", "finally", "serialize", "close", "checkpoint-stop", "end"} + } + assert len(worker_events) == 1 and threading.get_ident() not in worker_events + assert output["Status"] == ( + "FAILED" if result_kind in ("bad-json", "large-error") else "SUCCEEDED" + ) + + +@pytest.mark.parametrize("shape", ["method", "property", "dynamic"]) +def test_removed_handler_scope_api_is_not_inspected( + monkeypatch: pytest.MonkeyPatch, shape: str +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + calls = [] + + def legacy(*_args: Any) -> Any: + calls.append("unexpected") + raise AssertionError("Removed API called") + + choices: dict[str, dict[str, Any]] = { + "method": {"handler_context": legacy}, + "property": {"handler_context": property(legacy)}, + "dynamic": { + "__getattr__": lambda self, name: legacy() + if name == "handler_context" + else (_ for _ in ()).throw(AttributeError(name)) + }, + } + plugin = type("LegacyPlugin", (DurableInstrumentationPlugin,), choices[shape])() + handler = durable_execution(lambda _e, _c: "ok", plugins=[plugin]) + event, context = invocation() + assert handler(event, context)["Status"] == "SUCCEEDED" + assert calls == [] + + +@pytest.mark.parametrize( + "checkpoint_path", ["step-start", "large-result", "large-error"] +) +@pytest.mark.parametrize("retryable", [False, True]) +def test_checkpoint_failure_reports_prepared_outcome_after_worker_cleanup( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + checkpoint_path: str, + retryable: bool, +) -> None: + import aws_durable_execution_sdk_python.execution as execution + + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + if checkpoint_path != "step-start": + monkeypatch.setattr(execution, "LAMBDA_RESPONSE_SIZE_LIMIT", 128) + marker = contextvars.ContextVar("checkpoint-worker", default="host") + failure = CheckpointError( + "actual checkpoint failure", + error_category=( + CheckpointErrorCategory.INVOCATION + if retryable + else CheckpointErrorCategory.EXECUTION + ), + ) + events: list[tuple[str, int]] = [] + ends: list[tuple[InvocationEndInfo, str]] = [] + real_checkpoint = ExecutionState.checkpoint_batches_forever + real_close = ExecutionState.close + + def record(name: str) -> None: + events.append((name, threading.get_ident())) + + def checkpoint(state: ExecutionState) -> None: + record("checkpoint-start") + try: + real_checkpoint(state) + finally: + record("checkpoint-end") + + def close(state: ExecutionState) -> None: + record("close") + assert marker.get() == "plugin" + real_close(state) + record("closed") + + def service_checkpoint(*_args: Any, **_kwargs: Any) -> Any: + record("service-failure") + raise failure + + monkeypatch.setattr(ExecutionState, "checkpoint_batches_forever", checkpoint) + monkeypatch.setattr(ExecutionState, "close", close) + + class Plugin(DurableInstrumentationPlugin): + token: contextvars.Token[str] | None = None + + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + record("start") + self.token = marker.set("plugin") + + def on_invocation_end(self, info: InvocationEndInfo) -> None: + record("end") + ends.append((info, marker.get())) + assert self.token is not None + marker.reset(self.token) + self.token = None + + def body(_event: Any, ctx: DurableContext) -> str: + record("body") + try: + if checkpoint_path == "step-start": + ctx.step(lambda _step: "ok", name="failed-checkpoint") + if checkpoint_path == "large-error": + raise ValueError("x" * 256) + return "x" * 256 + finally: + record("finally") + + plugin = Plugin() + handler = durable_execution(body, plugins=[plugin]) + client = Mock(spec=DurableServiceClient) + client.checkpoint.side_effect = service_checkpoint + event, context = invocation(client) + for _ in range(2): + events.clear() + ends.clear() + if retryable: + with pytest.raises(CheckpointError) as caught: + handler(event, context) + assert caught.value is failure + else: + output = handler(event, context) + assert output["Status"] == "FAILED" + assert output["Error"]["ErrorMessage"] == str(failure) + assert output["Error"]["ErrorType"].endswith(".CheckpointError") + record("caller") + assert marker.get() == "host" and plugin.token is None + assert len(ends) == 1 + info, active = ends[0] + assert active == "plugin" + assert info.status is ( + InvocationStatus.RETRY if retryable else InvocationStatus.FAILED + ) + assert info.error is not None and info.error.message == str(failure) + assert info.error.type is not None and info.error.type.endswith( + ".CheckpointError" + ) + names = [name for name, _tid in events] + assert ( + names.index("start") + < names.index("checkpoint-start") + < names.index("service-failure") + ) + assert ( + names.index("body") + < names.index("finally") + < names.index("close") + < names.index("closed") + < names.index("end") + < names.index("caller") + ) + assert ( + names.index("service-failure") + < names.index("checkpoint-end") + < names.index("end") + ) + workers = { + tid + for name, tid in events + if name in {"start", "body", "finally", "close", "closed", "end"} + } + assert len(workers) == 1 and threading.get_ident() not in workers + assert not any( + record.exc_info and record.name == "aws_durable_execution_sdk_python.plugin" + for record in caplog.records + ) diff --git a/packages/aws-durable-execution-sdk-python/tests/operation/map_test.py b/packages/aws-durable-execution-sdk-python/tests/operation/map_test.py index 2670e6747..926b1dd0b 100644 --- a/packages/aws-durable-execution-sdk-python/tests/operation/map_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/operation/map_test.py @@ -7,6 +7,7 @@ import pytest # Mock the executor.execute method +from aws_durable_execution_sdk_python.plugin import PluginExecutor from aws_durable_execution_sdk_python.concurrency.models import ( BatchItem, BatchItemStatus, @@ -184,6 +185,8 @@ def mock_run_in_child_context(func, name, config): # Create a minimal ExecutionState mock class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -222,6 +225,8 @@ def mock_run_in_child_context(func, name, config): return func("mock_context") class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -354,6 +359,8 @@ def callable_func(ctx, item, idx, items): ) as mock_execute: class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -407,6 +414,8 @@ def callable_func(ctx, item, idx, items): executor_context.create_child_context = lambda *args, **kwargs: Mock() class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -461,6 +470,8 @@ def callable_func(ctx, item, idx, items): executor_context.create_child_context = lambda *args, **kwargs: child_context class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -510,6 +521,8 @@ def mock_summary_generator(result): executor_context.create_child_context = Mock(return_value=_child_ctx) class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -550,6 +563,8 @@ def callable_func(ctx, item, idx, items): executor_context.create_child_context = Mock(return_value=Mock()) # SLF001 class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -589,6 +604,8 @@ def func(ctx, item, index, array): config = MapConfig(summary_generator=None) class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -634,6 +651,8 @@ def callable_func(ctx, item, idx, items): # Mock execution state that indicates operation already succeeded class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -703,6 +722,8 @@ def callable_func(ctx, item, idx, items): # Mock execution state that indicates operation succeeded but children need replay class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -802,6 +823,8 @@ def test_func(ctx, item, idx, items): execution_count = 0 class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -1032,6 +1055,8 @@ def func(ctx, item, idx, items): return {"item": item.upper(), "index": idx} class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -1481,6 +1506,8 @@ def test_map_handler_defaults_summary_generator_for_user_config(): """ class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -1556,6 +1583,8 @@ def predicate(s: CompletionStatus) -> CompletionDecision: ) as mock_execute: class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass diff --git a/packages/aws-durable-execution-sdk-python/tests/operation/parallel_test.py b/packages/aws-durable-execution-sdk-python/tests/operation/parallel_test.py index 12231cfab..87d4d5a89 100644 --- a/packages/aws-durable-execution-sdk-python/tests/operation/parallel_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/operation/parallel_test.py @@ -8,6 +8,7 @@ import pytest +from aws_durable_execution_sdk_python.plugin import PluginExecutor from aws_durable_execution_sdk_python.concurrency.executor import ConcurrentExecutor from aws_durable_execution_sdk_python.identifier import OperationIdNamespace @@ -211,6 +212,8 @@ def func2(ctx): config = ParallelConfig(max_concurrency=2) class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -255,6 +258,8 @@ def func1(ctx): callables = [func1] class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -299,6 +304,8 @@ def func1(ctx): config = ParallelConfig(max_concurrency=5) class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -351,6 +358,8 @@ def func1(ctx): callables = [func1] class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -466,6 +475,8 @@ def func1(ctx): callables = [func1] class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -513,6 +524,8 @@ def mock_summary_generator(result): config = ParallelConfig(summary_generator=mock_summary_generator) class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -556,6 +569,8 @@ def func2(ctx): callables = [func1, func2] class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -606,6 +621,8 @@ def func3(ctx): config = ParallelConfig(summary_generator=None) class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -654,6 +671,8 @@ def func2(ctx): # Mock execution state that indicates operation already succeeded class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -723,6 +742,8 @@ def func1(ctx): # Mock execution state that indicates operation succeeded but children need replay class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -824,6 +845,8 @@ def task2(ctx): execution_count = 0 class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -1052,6 +1075,8 @@ def func3(ctx): callables = [func1, func2, func3] class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -1474,6 +1499,8 @@ def test_parallel_handler_defaults_summary_generator_for_user_config(): """A user config without a summary generator gets the default (JS parity).""" class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass @@ -1551,6 +1578,8 @@ def predicate(s: CompletionStatus) -> CompletionDecision: ) as mock_execute: class MockExecutionState: + _plugin_executor = PluginExecutor([]) + def register_branch_pool(self, pool): pass diff --git a/packages/aws-durable-execution-sdk-python/tests/plugin_test.py b/packages/aws-durable-execution-sdk-python/tests/plugin_test.py index 69cdfc502..0777fce2a 100644 --- a/packages/aws-durable-execution-sdk-python/tests/plugin_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/plugin_test.py @@ -2,6 +2,7 @@ import logging import pickle import unittest +from concurrent.futures import ThreadPoolExecutor from copy import deepcopy from dataclasses import asdict, fields from unittest.mock import MagicMock, patch @@ -1537,6 +1538,65 @@ def test_terminal_status_without_step_details_fires_operation_only(self): self.assertIn("operation_end:op-1", self.plugin.calls) + def test_checkpoint_does_not_repeat_an_observed_terminal_update(self): + """A resumed external result can also appear in the next checkpoint.""" + for status in ( + OperationStatus.SUCCEEDED, + OperationStatus.FAILED, + OperationStatus.CANCELLED, + OperationStatus.TIMED_OUT, + OperationStatus.STOPPED, + ): + with self.subTest(status=status): + plugin = _TrackingPlugin() + executor = PluginExecutor(plugins=[plugin]) + operation = self._make_operation(status=status) + with executor.run(): + # First completion delivered through UpdatedOperationIds. + executor.on_operation_update(operation) + # The next response carries the already observed state. + executor.on_operation_update( + [operation], + operations={operation.operation_id: operation}, + previous_operations={operation.operation_id: operation}, + ) + self.assertEqual(plugin.calls.count("operation_end:op-1"), 1) + + def test_checkpoint_terminal_transition_still_emits(self): + previous = self._make_operation(status=OperationStatus.STARTED) + operation = self._make_operation(status=OperationStatus.SUCCEEDED) + with self.executor.run(): + self.executor.on_operation_update( + [operation], + operations={operation.operation_id: operation}, + previous_operations={previous.operation_id: previous}, + ) + self.assertEqual(self.plugin.calls.count("operation_end:op-1"), 1) + + def test_checkpoint_preserves_first_terminal_notification(self): + """State can predate notification, e.g. without UpdatedOperationIds.""" + operation = self._make_operation(status=OperationStatus.SUCCEEDED) + with self.executor.run(): + self.executor.on_operation_update( + [operation], + operations={operation.operation_id: operation}, + previous_operations={operation.operation_id: operation}, + ) + self.assertEqual(self.plugin.calls.count("operation_end:op-1"), 1) + + def test_terminal_notification_resets_between_invocations(self): + operation = self._make_operation(status=OperationStatus.SUCCEEDED) + for _ in range(2): + with self.executor.run(): + self.executor.on_operation_update(operation) + self.assertEqual(self.plugin.calls.count("operation_end:op-1"), 2) + + def test_concurrent_completion_notifications_emit_once(self): + operation = self._make_operation(status=OperationStatus.SUCCEEDED) + with self.executor.run(), ThreadPoolExecutor(max_workers=8) as workers: + list(workers.map(self.executor.on_operation_update, [operation] * 16)) + self.assertEqual(self.plugin.calls.count("operation_end:op-1"), 1) + def test_non_terminal_status_without_step_details_fires_nothing(self): op = self._make_operation(status=OperationStatus.STARTED, step_details=None) diff --git a/pyproject.toml b/pyproject.toml index 5291b0c5e..80eb21e0c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -131,44 +131,48 @@ dependencies = [ [tool.hatch.envs.dev-examples.scripts] test = "pytest packages/aws-durable-execution-sdk-python-examples/test {args}" -[tool.hatch.envs.test-pypi-otel] -# Test new exclusivity capability against an installed capable core. -# Legacy valid-registration combinations are verified separately against 2.0.x. -# Override inherited workspace membership: core must come from an installed distribution. -workspace.members = ["packages/aws-durable-execution-sdk-python-otel"] +# Install immutable artifacts built from this checkout, including before the +# first compatible core is published. No editable workspace packages in this lane. +[tool.hatch.envs.test-wheel-otel] +detached = true +workspace.members = [] dependencies = [ "aws-durable-execution-sdk-python-testing>=1.2.1", - "aws-durable-execution-sdk-python>=2.1.0", "opentelemetry-sdk>=1.20.0", "opentelemetry-propagator-aws-xray", "pytest", - "pytest-cov", - "coverage[toml]", -] -pre-install-commands = [ - "pip install -e packages/aws-durable-execution-sdk-python-otel", + "packaging", ] -[tool.hatch.envs.test-pypi-otel.scripts] -test = "pytest packages/aws-durable-execution-sdk-python-otel/tests {args}" +[tool.hatch.envs.test-wheel-otel.scripts] +test = [ + "python .github/scripts/install_otel_test_wheels.py", + "pytest packages/aws-durable-execution-sdk-python-otel/tests {args}", + "pytest .github/tests/otel_lifecycle_compatibility_test.py {args}", +] -[tool.hatch.envs.test-pypi-otel-legacy] -template = "test-pypi-otel" -workspace.members = ["packages/aws-durable-execution-sdk-python-otel"] +# A core-only upgrade retains the released plugin's more limited Invocation +# fallback. Validate the real published plugin separately from the new pair. +[tool.hatch.envs.test-wheel-otel-legacy] +template = "test-wheel-otel" +detached = true +workspace.members = [] dependencies = [ - "aws-durable-execution-sdk-python>=2.0.0,<2.1.0", + "aws-durable-execution-sdk-python-otel==1.0.0", "aws-durable-execution-sdk-python-testing>=1.2.1", "opentelemetry-sdk>=1.20.0", "opentelemetry-propagator-aws-xray", "pytest", - "pytest-cov", - "coverage[toml]", + "packaging", ] -[tool.hatch.envs.test-pypi-otel-legacy.scripts] +[tool.hatch.envs.test-wheel-otel-legacy.env-vars] +OTEL_COMPAT_LEGACY = "1" + +[tool.hatch.envs.test-wheel-otel-legacy.scripts] test = [ - "python -c 'from pathlib import Path; from importlib.metadata import version; from packaging.version import Version; import aws_durable_execution_sdk_python.execution as core; assert Version(version(\"aws-durable-execution-sdk-python\")).release[:2] == (2, 0); assert \"site-packages\" in Path(core.__file__).resolve().parts; print(version(\"aws-durable-execution-sdk-python\"), core.__file__)' ", - "pytest packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_view_registration.py -k 'one_view_and_unrelated_plugin_suspend_resume or no_otel_plugin_remains_valid or execution_constructor_retains_ambient_log_correlation' {args}" + "python .github/scripts/install_otel_test_wheels.py --legacy-plugin", + "pytest .github/tests/otel_lifecycle_compatibility_test.py {args}", ] [tool.hatch.envs.test-pypi-examples]