diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 63c4b44c6..174a998bd 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -60,6 +60,19 @@ jobs: run: hatch run types:check - name: Run tests + coverage run: hatch run test:cov + - name: Test OTel with the minimum supported core + run: | + hatch run test-pypi-otel:python - <<'PYTHON' + from importlib.metadata import version + from pathlib import Path + import aws_durable_execution_sdk_python.plugin as plugin + + assert version("aws-durable-execution-sdk-python") == "2.0.0" + assert "site-packages" in Path(plugin.__file__).parts + assert not hasattr(plugin.DurableInstrumentationPlugin, "handler_context") + assert not hasattr(plugin, "DURABLE_INSTRUMENTATION_HANDLER_CONTEXT_API_VERSION") + PYTHON + hatch run test-pypi-otel:test - name: Build distribution run: | for pkg in packages/*/; do diff --git a/.github/workflows/cloud-tests.yml b/.github/workflows/cloud-tests.yml index 3d3b69950..fc4b0f57e 100644 --- a/.github/workflows/cloud-tests.yml +++ b/.github/workflows/cloud-tests.yml @@ -114,7 +114,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/CONTRIBUTING.md b/CONTRIBUTING.md index ca3ff1008..3e46a51ae 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -75,10 +75,14 @@ hatch run dev-examples:test # run examples tests only To verify packages work against the published PyPI version of the core SDK (rather than the local workspace): ```bash -hatch run test-pypi-otel:test # test otel against PyPI core SDK +hatch run test-pypi-otel:test # test otel against the minimum supported core (2.0.0) hatch run test-pypi-examples:test # test examples against PyPI core SDK ``` +The OTel PyPI environment excludes the local core and pins the minimum supported +release, so newer PyPI releases cannot remove legacy compatibility coverage. Use +`hatch run dev-otel:test` for the current workspace core. + ### Package-level commands Some commands still run from within a package directory: 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 5e78f2a9c..b99ec4dd7 100644 --- a/packages/aws-durable-execution-sdk-python-otel/README.md +++ b/packages/aws-durable-execution-sdk-python-otel/README.md @@ -156,6 +156,27 @@ lambda_.Function( ) ``` +### Handler context propagation + +A core SDK with handler-worker context propagation carries the context established +by invocation-start hooks into the handler. Invocation view preserves an active +ambient span on the canonical execution trace; when that context is absent or +belongs to a different trace, its optional `handler_context` scope makes the Invocation +span current only while the handler runs. The scope closes on the same worker in +reverse plugin order, including on failure and suspension, without changing the +invocation-hook caller. Older cores ignore this optional scope and retain their +existing behavior; install the updated core as well to get handler context propagation. +The bundled OTel classes explicitly opt in with `__durable_handler_context_api__ = 1`. +Custom subclasses must repeat that literal marker on their own concrete class to +use the scope; inherited or instance markers are ignored. Unopted legacy helpers +and properties with the same name are never inspected. Older cores ignore the +marker without importing any new core API. +Execution view similarly restores the Workflow span inside the handler scope if +another invocation-start hook clears the active span or switches to an unrelated +trace. Both views retain valid same-trace parents and baggage, and restore the +worker's previous context when the scope ends. +Existing plugin registration, factory lifetime, and checkpoint formats are unchanged. + ### 3. In your Lambda handler (index.py) ```python @@ -297,6 +318,21 @@ The resolved decision is applied to Workflow, Invocation, operation, and attempt spans. This avoids independently querying stateful or ratio-based samplers for each durable span in the same invocation. +Invocation hooks retain their caller thread and registration order. With the +updated core, invocation-local context-variable bindings are isolated from the +host: hooks see the incoming context and the handler receives their resulting +context, while invocation exit restores the host's original bindings even if a +plugin fails during setup or cleanup. Plugins must not use invocation context +bindings to mutate the host context after the invocation has returned. Older +supported cores retain their existing lifecycle behavior, including the +execution-view limitation when later plugins open invocation context scopes. +The new isolation applies only when plugins are registered. If an invocation-start +hook or handler-scope entry raises, subsequent setup and the handler retain the +bindings from before that hook. Successful scopes still clean up in their original +context, preserving token ownership. This isolates context-variable bindings; +it does not undo a plugin's mutations to shared objects or external side effects. + + ### Log Correlation When `enrich_logger=True` (the default), the plugin installs a logging filter on diff --git a/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/execution_plugin.py b/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/execution_plugin.py index 1a46cafcb..c6bc2584a 100644 --- a/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/execution_plugin.py +++ b/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/execution_plugin.py @@ -31,9 +31,11 @@ from __future__ import annotations +import contextlib import datetime import logging import threading +from collections.abc import Iterator from typing import Any from aws_durable_execution_sdk_python.plugin import ( @@ -123,6 +125,8 @@ class ExecutionOtelPlugin(DurableInstrumentationPlugin): span). """ + __durable_handler_context_api__ = 1 + def __init__(self, config: OtelPluginConfig | None = None) -> None: self._config = config or OtelPluginConfig() self._context_extractor: ContextExtractor = ( @@ -487,6 +491,24 @@ def on_invocation_start(self, info: InvocationStartInfo) -> None: ), ) + @contextlib.contextmanager + def handler_context(self, info: InvocationStartInfo) -> Iterator[None]: + """Keep handler instrumentation on this execution's trace in its worker.""" + ambient = trace.get_current_span().get_span_context() + workflow = self._workflow_span + token = None + if ( + self._tracing_enabled + and workflow is not None + and (not ambient.is_valid or ambient.trace_id != self._execution_trace_id) + ): + token = otel_context.attach(trace.set_span_in_context(workflow)) + try: + yield + finally: + if token is not None: + otel_context.detach(token) + def _start_workflow_span(self, info: InvocationStartInfo) -> None: """Install a non-recording placeholder for the execution-scoped Workflow span. 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 012f17f2f..980764f3e 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 @@ -2,9 +2,11 @@ from __future__ import annotations +import contextlib import datetime import logging import threading +from collections.abc import Iterator from typing import Any from aws_durable_execution_sdk_python.plugin import ( @@ -104,6 +106,8 @@ class InvocationOtelPlugin(DurableInstrumentationPlugin): provider installed by the ADOT Lambda layer). """ + __durable_handler_context_api__ = 1 + DEFAULT_INSTRUMENT_NAME = "aws-durable-execution-sdk-python" def __init__(self, config: OtelPluginConfig | None = None) -> None: @@ -286,12 +290,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 outside the optional handler-worker context scope. Returns: A valid SpanContext, or None if no span is active. @@ -579,6 +579,24 @@ def on_invocation_start(self, info: InvocationStartInfo) -> None: attributes=self._extract_attributes(info), ) + @contextlib.contextmanager + def handler_context(self, info: InvocationStartInfo) -> Iterator[None]: + """Bind the fallback only inside the SDK-owned handler worker scope.""" + ambient = trace.get_current_span().get_span_context() + invocation_span = self._get_span(None) + token = None + if ( + self._tracing_enabled + and invocation_span is not None + and (not ambient.is_valid or ambient.trace_id != self._execution_trace_id) + ): + token = context.attach(trace.set_span_in_context(invocation_span)) + try: + yield + finally: + if token is not None: + context.detach(token) + def _start_workflow_span(self, info: InvocationStartInfo) -> None: """Install a non-recording placeholder for the execution-scoped Workflow span. @@ -655,6 +673,10 @@ def on_invocation_end(self, info: InvocationEndInfo) -> None: self._reset_state() return + # User execution has finished; the worker has already closed its handler + # context scope without modifying the invocation-hook caller. + 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: 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..427854750 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,17 @@ OperationType, StepDetails, ) +from aws_durable_execution_sdk_python import plugin as core_plugin_api +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 +259,239 @@ 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 + + +@pytest.mark.parametrize( + ("plugin_type", "extra_context_plugin"), + [ + (InvocationOtelPlugin, False), + (InvocationOtelPlugin, True), + (ExecutionOtelPlugin, False), + ] + # Execution-view caller isolation requires the coordinated newer core. + # Released 2.0.x retains the pre-existing same-order teardown limitation; + # the legacy lane continues checking its supported combinations above. + + ( + [(ExecutionOtelPlugin, True)] + + [(ExecutionOtelPlugin, kind) for kind in ("same", "unrelated", "absent")] + if getattr( + core_plugin_api, "DURABLE_INSTRUMENTATION_HANDLER_CONTEXT_API_VERSION", None + ) + == 1 + else [] + ), +) +@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) + # The documented PyPI compatibility environment deliberately uses an older + # core. Keep exercising its supported operation tracing and lifecycle while + # asserting the new handler contract only when that core exposes the scope. + supports_handler_context = ( + getattr( + core_plugin_api, "DURABLE_INSTRUMENTATION_HANDLER_CONTEXT_API_VERSION", None + ) + == 1 + ) + 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" if supports_handler_context else None + ) + 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 + + # An unrelated plugin may own a caller-thread OTel baggage scope. The + # invocation-view fallback must never become part of its saved token. + 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 = ( + [None, None] + if not supports_handler_context + else [ + 0xCAFE + if extra_context_plugin == "same" and not reverse_plugins + else derive_workflow_span_id(EXECUTION_ARN) + ] + * 2 + if plugin_type is ExecutionOtelPlugin + else 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 + ] + ) + 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 + and ( + (span.context.trace_id == XRAY_TRACE_ID) == supports_handler_context + ) + 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/test_execution_plugin.py b/packages/aws-durable-execution-sdk-python-otel/tests/test_execution_plugin.py index 5c3cc3eca..f9043c816 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 ( @@ -1620,3 +1621,73 @@ 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_handler_scope_keeps_execution_trace_and_restores_context( + 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, + ) + ) + info = _invocation_start_info() + plugin.on_invocation_start(info) + workflow = trace.get_current_span().get_span_context() + assert workflow.is_valid + assert workflow.span_id == derive_workflow_span_id(EXECUTION_ARN) + caller = baggage.set_baggage("customer", "preserved", Context()) + expected = workflow + if ambient_kind != "absent": + ambient = SpanContext( + trace_id=workflow.trace_id if ambient_kind == "same" else 1, + span_id=0xCAFE, + is_remote=False, + trace_flags=workflow.trace_flags, + ) + caller = trace.set_span_in_context(NonRecordingSpan(ambient), caller) + if ambient_kind == "same": + expected = ambient + token = otel_context.attach(caller) + error = ValueError("user failure") + try: + + def body() -> None: + with plugin.handler_context(info): + assert trace.get_current_span().get_span_context() == expected + assert baggage.get_baggage("customer") == "preserved" + if raises: + raise error + + if raises: + with pytest.raises(ValueError) as caught: + body() + assert caught.value is error + else: + body() + assert otel_context.get_current() is caller + finally: + otel_context.detach(token) + plugin.on_invocation_end(_invocation_end_info()) + provider.shutdown() + + +@pytest.mark.parametrize("completed", [False, True]) +def test_handler_scope_without_live_workflow_is_noop(completed: bool) -> None: + plugin, _ = _create_plugin() + info = _invocation_start_info() + if completed: + plugin.on_invocation_start(info) + plugin.on_invocation_end(_invocation_end_info()) + caller = otel_context.get_current() + with plugin.handler_context(info): + assert otel_context.get_current() is caller + assert otel_context.get_current() is caller 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 f108ef728..7e8fb2065 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 @@ -2198,3 +2198,71 @@ def test_nested_suspension_unwinds_scopes_in_reverse_order(): assert plugin._context_tokens == {} 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_handler_context_preserves_or_replaces_parent_and_restores_baggage( + ambient_kind: str, raises: bool, sampler: Sampler +) -> None: + plugin, _ = _create_plugin_with_sampler(sampler) + info = _invocation_start_info() + plugin.on_invocation_start(info) + invocation = plugin.get_current_span_context() + assert invocation is not None and invocation.is_valid + caller = baggage.set_baggage("tenant", "scope-test", Context()) + expected = invocation + if ambient_kind != "absent": + ambient = SpanContext( + trace_id=( + invocation.trace_id + if ambient_kind == "same" + else (1 if invocation.trace_id != 1 else 2) + ), + span_id=0x42, + is_remote=False, + trace_flags=invocation.trace_flags, + ) + caller = trace.set_span_in_context(NonRecordingSpan(ambient), caller) + if ambient_kind == "same": + expected = ambient + token = otel_context.attach(caller) + error = ValueError("handler error") + try: + + def body() -> None: + with plugin.handler_context(info): + active = trace.get_current_span().get_span_context() + assert active == expected + assert active.is_valid + assert baggage.get_baggage("tenant") == "scope-test" + if raises: + raise error + + if raises: + with pytest.raises(ValueError) as caught: + body() + assert caught.value is error + else: + body() + assert otel_context.get_current() is caller + assert baggage.get_baggage("tenant") == "scope-test" + finally: + otel_context.detach(token) + plugin.on_invocation_end(_invocation_end_info()) + + +@pytest.mark.parametrize("completed", [False, True]) +def test_handler_context_without_live_invocation_leaves_context_unchanged( + completed: bool, +) -> None: + plugin, _ = _create_plugin() + info = _invocation_start_info() + if completed: + plugin.on_invocation_start(info) + plugin.on_invocation_end(_invocation_end_info()) + caller = otel_context.get_current() + with plugin.handler_context(info): + assert otel_context.get_current() is caller + assert otel_context.get_current() is caller 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 22b416384..cc4c01ee9 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 @@ -74,9 +74,9 @@ 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: +def test_pypi_compatibility_environment_pins_minimum_supported_core() -> None: dependencies = _load_pyproject(REPOSITORY_ROOT / "pyproject.toml")["tool"]["hatch"][ "envs" ]["test-pypi-otel"]["dependencies"] - assert CORE_DEPENDENCY in dependencies + assert "aws-durable-execution-sdk-python==2.0.0" in dependencies diff --git a/packages/aws-durable-execution-sdk-python/README.md b/packages/aws-durable-execution-sdk-python/README.md index bf7776f45..bfaf92157 100644 --- a/packages/aws-durable-execution-sdk-python/README.md +++ b/packages/aws-durable-execution-sdk-python/README.md @@ -77,6 +77,25 @@ Provider names must be unique across installed distributions. Missing, ambiguous, incompatible, or invalid providers raise `PluginLoadError` during handler initialization with the provider and distribution details. +### Optional handler context scopes + +A plugin can declare `__durable_handler_context_api__ = 1` directly on its +concrete class and implement `handler_context(info)` returning a context manager. +The updated core enters these scopes around the top-level handler on its worker +thread, in registration order, and closes them in reverse order. Cleanup receives +no handler exception and cannot suppress or replace its outcome. Invocation hooks +retain their original thread and order; failed setup bindings are discarded and +successful scope cleanup stays in the context that owns its tokens. + +The marker must be the literal integer `1`; instance and inherited markers do not +opt in. A subclass must redeclare the marker to adopt this new hook. Unopted legacy +helpers, properties and dynamic attributes named `handler_context` are untouched. +The generic plugin base supplies neither a marker nor a default method. Core +support is advertised by module constant +`aws_durable_execution_sdk_python.plugin.DURABLE_INSTRUMENTATION_HANDLER_CONTEXT_API_VERSION`. +Older cores ignore this optional API. The provider API version and dependency +requirements are unchanged. + ## 🚀 Quick Start Install the execution SDK: 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..732f6d6e8 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,6 +1,7 @@ from __future__ import annotations import contextlib +import contextvars import functools import json import logging @@ -188,7 +189,8 @@ 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]: @@ -313,7 +315,20 @@ def wrapper(event: Any, context: LambdaContext) -> MutableMapping[str, Any]: logger.debug( "%s entering user-space...", invocation_input.durable_execution_arn ) - user_future = executor.submit(func, input_event, durable_context) + if configured_plugins: + # Invocation-start hooks can establish tracing and other contextvars. + # Context.run restores worker bindings on both return and failure. + user_future = executor.submit( + contextvars.copy_context().run, + plugin_executor.run_handler, + func, + input_event, + durable_context, + ) + else: + # Preserve the original fresh-worker context for uninstrumented + # handlers, including the absence of caller ContextVar bindings. + user_future = executor.submit(func, input_event, durable_context) logger.debug( "%s waiting for user code completion...", 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..f9e202017 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,11 +1,12 @@ from __future__ import annotations import contextlib +import contextvars import copy import datetime import functools import logging -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from enum import Enum @@ -29,6 +30,7 @@ logger = logging.getLogger(__name__) DURABLE_INSTRUMENTATION_PLUGIN_API_VERSION = 1 +DURABLE_INSTRUMENTATION_HANDLER_CONTEXT_API_VERSION = 1 class InvocationStatus(Enum): @@ -460,12 +462,26 @@ class DurableInstrumentationPluginProvider: plugin_api_version: int +def _handler_context_api_enabled( + plugin_type: type[DurableInstrumentationPlugin], +) -> bool: + """Read only the concrete class namespace, bypassing metaclass descriptors.""" + namespace = type.__dict__["__dict__"].__get__(plugin_type, type(plugin_type)) + version = namespace.get("__durable_handler_context_api__") + return ( + type(version) is int + and version == DURABLE_INSTRUMENTATION_HANDLER_CONTEXT_API_VERSION + ) + + class PluginExecutor: def __init__(self, plugins: list[DurableInstrumentationPlugin] | None): self._plugins = plugins or [] self._executor: ThreadPoolExecutor | None = None self._invocation_status: InvocationStartInfo | None = None self._operations_provider: Callable[[], Mapping[str, Operation]] | None = None + self._startup_context: contextvars.Context | None = None + self._invocation_contexts: list[contextvars.Context | None] = [] @contextlib.contextmanager def run(self): @@ -479,12 +495,14 @@ 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) @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 +525,105 @@ 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) + @contextlib.contextmanager + def _safe_handler_context( + self, plugin: DurableInstrumentationPlugin, info: InvocationStartInfo + ) -> Iterator[bool]: + # Old plugin objects may not inherit this core's new optional method. + scope = None + succeeded = True + try: + # Old plugins may have an unrelated helper/property with this name. + # Never even inspect it unless this concrete class explicitly opts in. + factory = ( + getattr(plugin, "handler_context", None) + if _handler_context_api_enabled(type(plugin)) + else None + ) + if factory is not None: + scope = factory(info) + scope.__enter__() + except Exception: + scope = None + succeeded = False + logger.exception( + "Plugin %s handler context failed", plugin.__class__.__name__ + ) + try: + yield succeeded + finally: + if scope is not None: + try: + scope.__exit__(None, None, None) + except Exception: + logger.exception( + "Plugin %s handler context cleanup failed", + plugin.__class__.__name__, + ) + + def run_handler(self, handler: Callable[..., Any], *args: Any) -> Any: + """Run the user handler inside optional, balanced plugin context scopes.""" + if not self._plugins or self._invocation_status is None: + return handler(*args) + owner = ( + self._startup_context.copy() + if self._startup_context is not None + else contextvars.copy_context() + ) + with contextlib.ExitStack() as scopes: + for plugin in self._plugins: + before = owner.copy() + scope = self._safe_handler_context(plugin, self._invocation_status) + succeeded = owner.run(scope.__enter__) + # Keep finalizers with their entry Context: ContextVar tokens + # cannot be reset in a copy, even when its bindings are identical. + scopes.callback(owner.run, scope.__exit__, None, None, None) + if not succeeded: + owner = before + return owner.run(handler, *args) + def _snapshot_operation_infos( self, operations_provider: Callable[[], Mapping[str, Operation]] | None, @@ -840,8 +945,7 @@ def _is_terminal_status(status): @property def handle_durable_output(self): def decorator(func: Callable[[Any, LambdaContext], MutableMapping[str, Any]]): - @functools.wraps(func) - def wrapper(event: Any, context: LambdaContext): + def invoke(event: Any, context: LambdaContext): with self.run(): try: output = func(event, context) @@ -858,6 +962,16 @@ def wrapper(event: Any, context: LambdaContext): ) raise + @functools.wraps(func) + def wrapper(event: Any, context: LambdaContext): + if not self._plugins: + return invoke(event, context) + # Keep hooks on their existing caller thread and in registration + # order, but isolate their context bindings from the host. Two + # plugins can otherwise restore a stale predecessor when their + # invocation-end hooks close scopes in the original order. + return contextvars.copy_context().run(invoke, event, context) + return wrapper return decorator 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..215624ea3 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/handler_worker_context_test.py @@ -0,0 +1,562 @@ +"""Focused handler-dispatch tests with a mock service and real worker thread.""" + +from __future__ import annotations + +import contextvars +from collections.abc import Callable +from concurrent.futures import Future, 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 InvocationError +from aws_durable_execution_sdk_python.execution import durable_execution +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationEndInfo, + InvocationStartInfo, + InvocationStatus, +) + + +@pytest.mark.parametrize("outcome", ["success", "failure", "retry"]) +@pytest.mark.parametrize("plugin_mode", ["none", "healthy", "partial-failure"]) +def test_handler_worker_preserves_context_and_restores_its_caller( + monkeypatch: pytest.MonkeyPatch, outcome: str, plugin_mode: str +) -> None: + marker = contextvars.ContextVar("handler-worker-context", default="worker-empty") + seen: list[str] = [] + worker_boundaries: list[tuple[str, str]] = [] + statuses: list[InvocationStatus] = [] + + class ClaimPlugin(DurableInstrumentationPlugin): + token: contextvars.Token[str] | None = None + + def on_invocation_start(self, info: InvocationStartInfo) -> None: + self.token = marker.set("invocation-start") + if plugin_mode == "partial-failure": + raise ValueError("partial plugin setup") + + 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 + + def body(_event: Any, _context: DurableContext) -> str: + seen.append(marker.get()) + marker.set("worker-mutation") + if outcome == "failure": + raise ValueError("handler failure") + if outcome == "retry": + raise InvocationError("handler retry") + return "ok" + + class ObservingExecutor(ThreadPoolExecutor): + def submit( + self, fn: Callable[..., Any], /, *args: Any, **kwargs: Any + ) -> Future[Any]: + if fn is not body and body not in args: + return super().submit(fn, *args, **kwargs) + + def observe() -> Any: + before = marker.get() + try: + return fn(*args, **kwargs) + finally: + # This runs outside Context.run, on the actual SDK worker. + worker_boundaries.append((before, marker.get())) + + return super().submit(observe) + + monkeypatch.setattr( + "aws_durable_execution_sdk_python.execution.ThreadPoolExecutor", + ObservingExecutor, + ) + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + client = Mock() + handler = durable_execution( + body, + boto3_client=client, + plugins=[ClaimPlugin()] if plugin_mode != "none" else [], + ) + event = { + "DurableExecutionArn": "test-arn/handler-context", + "CheckpointToken": "test-token", + "InitialExecutionState": { + "Operations": [ + { + "Id": "handler-context", + "Type": "EXECUTION", + "Status": "STARTED", + "ExecutionDetails": {"InputPayload": "{}"}, + } + ], + "NextMarker": "", + }, + } + lambda_context = Mock() + lambda_context.aws_request_id = "context-request" + lambda_context.client_context = None + lambda_context.identity = None + lambda_context._epoch_deadline_time_in_ms = 0 + lambda_context.invoked_function_arn = "test-arn" + lambda_context.tenant_id = None + token = marker.set("caller") + try: + if outcome == "retry": + with pytest.raises(InvocationError, match="handler retry"): + handler(event, lambda_context) + else: + result = handler(event, lambda_context) + assert result["Status"] == ( + "SUCCEEDED" if outcome == "success" else "FAILED" + ) + assert marker.get() == "caller" + finally: + marker.reset(token) + assert len(worker_boundaries) == 1 + worker_before, worker_after = worker_boundaries[0] + if plugin_mode != "none": + assert seen == [ + "caller" if plugin_mode == "partial-failure" else "invocation-start" + ] + assert worker_after == worker_before + assert worker_after != "worker-mutation" + assert statuses == [ + { + "success": InvocationStatus.SUCCEEDED, + "failure": InvocationStatus.FAILED, + "retry": InvocationStatus.RETRY, + }[outcome] + ] + else: + # No plugin means the original direct worker call: caller bindings are + # absent, and mutations belong to the worker's own context. + assert seen == [worker_before] + assert seen != ["caller"] + assert worker_after == "worker-mutation" + assert statuses == [] + client.checkpoint_durable_execution.assert_not_called() + + +@pytest.mark.parametrize("failure", [None, "body", "enter", "exit"]) +def test_optional_handler_scopes_are_balanced_and_cannot_change_outcome( + failure: str | None, +) -> None: + from contextlib import contextmanager + from datetime import UTC, datetime + from collections.abc import Iterator + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + marker = contextvars.ContextVar("handler-scope", default="caller") + events: list[str] = [] + + class ScopePlugin(DurableInstrumentationPlugin): + __durable_handler_context_api__ = 1 + + def __init__(self, name: str): + self.name = name + + @contextmanager + def handler_context(self, info: InvocationStartInfo) -> Iterator[None]: + assert info.execution_arn == "handler-scope" + events.append("enter-" + self.name) + if self.name == "inner" and failure == "enter": + raise ValueError("plugin entry failure") + token = marker.set(self.name) + try: + yield + finally: + marker.reset(token) + events.append("exit-" + self.name) + if self.name == "inner" and failure == "exit": + raise ValueError("plugin cleanup failure") + + executor = PluginExecutor([ScopePlugin("outer"), ScopePlugin("inner")]) + + def body() -> str: + assert marker.get() == ("outer" if failure == "enter" else "inner") + events.append("body") + if failure in ("body", "exit"): + raise RuntimeError("original handler failure") + return "ok" + + with executor.run(): + executor.on_invocation_start("handler-scope", True, datetime.now(UTC), None) + if failure in ("body", "exit"): + with pytest.raises(RuntimeError, match="original handler failure"): + executor.run_handler(body) + else: + assert executor.run_handler(body) == "ok" + assert marker.get() == "caller" + assert events == ["enter-outer", "enter-inner", "body"] + ( + ["exit-outer"] if failure == "enter" else ["exit-inner", "exit-outer"] + ) + + +def test_handler_accepts_plugin_without_optional_scope() -> None: + from datetime import UTC, datetime + from types import SimpleNamespace + from typing import cast + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + legacy = cast( + DurableInstrumentationPlugin, + SimpleNamespace(on_invocation_start=lambda info: None), + ) + executor = PluginExecutor([legacy]) + with executor.run(): + executor.on_invocation_start("legacy", True, datetime.now(UTC), None) + assert executor.run_handler(lambda: "unchanged") == "unchanged" + + +@pytest.mark.parametrize("outcome", ["SUCCEEDED", "PENDING", "FAILED", "retry"]) +@pytest.mark.parametrize("reverse", [False, True]) +@pytest.mark.parametrize("hook_failure", [None, "start", "end"]) +def test_invocation_context_scopes_do_not_escape_to_host( + outcome: str, reverse: bool, hook_failure: str | None +) -> None: + """Legacy hook order must not leave an already-ended plugin scope current.""" + from datetime import UTC, datetime + import threading + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + marker = contextvars.ContextVar("invocation-scope", default="host") + events: list[tuple[str, str, str, int]] = [] + caller_thread = threading.get_ident() + + 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: + events.append(("start", self.name, marker.get(), threading.get_ident())) + self.token = marker.set(self.name) + if self.name == names[0] and hook_failure == "start": + raise ValueError("plugin initialization failed") + + def on_invocation_end(self, _info: InvocationEndInfo) -> None: + events.append(("end", self.name, marker.get(), threading.get_ident())) + if self.name == names[0] and hook_failure == "end": + raise ValueError("plugin finalization failed") + assert self.token is not None + marker.reset(self.token) + self.token = None + + names = ["first", "second"] + if reverse: + names.reverse() + executor = PluginExecutor([ScopePlugin(name) for name in names]) + handler_failure = InvocationError("retry") + expected_output = {"Status": outcome} + + @executor.handle_durable_output + def invoke(_event: Any, _context: Any) -> dict[str, str]: + executor.on_invocation_start("test", True, datetime.now(UTC), None) + assert executor.run_handler(marker.get) == names[-1] + if outcome == "retry": + raise handler_failure + return expected_output + + token = marker.set("incoming") + try: + for _ in range(2): + if outcome == "retry": + with pytest.raises(InvocationError, match="retry") as caught: + invoke({}, None) + assert caught.value is handler_failure + else: + assert invoke({}, None) is expected_output + assert marker.get() == "incoming" + finally: + marker.reset(token) + assert [(kind, name) for kind, name, _, _ in events] == [ + (kind, name) for _ in range(2) for kind in ("start", "end") for name in names + ] + assert all(thread == caller_thread for _, _, _, thread in events) + assert [ + value for kind, name, value, _ in events if kind == "start" and name == names[0] + ] == ["incoming", "incoming"] + + +def test_no_plugin_invocation_keeps_original_caller_context_semantics() -> None: + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + marker = contextvars.ContextVar("no-plugin-caller", default="host") + executor = PluginExecutor([]) + + @executor.handle_durable_output + def invoke(_event: Any, _context: Any) -> dict[str, str]: + marker.set("caller-side-change") + return {"Status": "SUCCEEDED"} + + token = marker.set("incoming") + try: + invoke({}, None) + assert marker.get() == "caller-side-change" + finally: + marker.reset(token) + + +@pytest.mark.parametrize("stage", ["start", "factory", "enter"]) +@pytest.mark.parametrize("bad_first", [False, True]) +@pytest.mark.parametrize("outcome", ["SUCCEEDED", "PENDING", "FAILED", "retry"]) +def test_failed_plugin_setup_discards_partial_context_bindings( + stage: str, bad_first: bool, outcome: str, caplog: pytest.LogCaptureFixture +) -> None: + """Keep healthy bindings, unset bindings, hook order, and token ownership.""" + from contextlib import nullcontext + from datetime import UTC, datetime + from typing import ContextManager + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + marker = contextvars.ContextVar("partial-setup", default="default") + new_binding = contextvars.ContextVar[str]("partial-setup-no-default") + events: list[tuple[str, str]] = [] + cleanup: list[str] = [] + healthy_inputs: list[tuple[str, str | None]] = [] + + class Scope: + def __init__(self, plugin: SetupPlugin) -> None: + self.plugin = plugin + + def __enter__(self) -> None: + self.plugin.bind("enter") + + def __exit__(self, *_args: Any) -> None: + self.plugin.reset("exit") + + class SetupPlugin(DurableInstrumentationPlugin): + __durable_handler_context_api__ = 1 + + def __init__(self, name: str) -> None: + self.name = name + self.token: contextvars.Token[str] | None = None + + def bind(self, where: str) -> None: + events.append((where, self.name)) + if self.name == "healthy": + healthy_inputs.append((marker.get(), new_binding.get(None))) + self.token = marker.set(self.name) + if self.name == "bad": + new_binding.set("partial") + raise ValueError("partial plugin setup") + + def reset(self, where: str) -> None: + events.append((where, self.name)) + assert self.token is not None + marker.reset(self.token) + self.token = None + cleanup.append(self.name) + + def on_invocation_start(self, _info: InvocationStartInfo) -> None: + if stage == "start": + self.bind("start") + + def on_invocation_end(self, _info: InvocationEndInfo) -> None: + if stage == "start": + self.reset("end") + + def handler_context(self, _info: InvocationStartInfo) -> ContextManager[None]: + if stage == "start": + return nullcontext() + if self.name == "bad" and stage == "factory": + self.bind("factory") + return Scope(self) + + names = ["bad", "healthy"] if bad_first else ["healthy", "bad"] + executor = PluginExecutor([SetupPlugin(name) for name in names]) + output = {"Status": outcome} + failure = InvocationError("original retry") + handler_failure = RuntimeError("original handler error") + + def body() -> dict[str, str]: + assert marker.get() == "healthy" + with pytest.raises(LookupError): + new_binding.get() + if outcome == "retry": + raise failure + if outcome == "FAILED": + raise handler_failure + return output + + @executor.handle_durable_output + def invoke(_event: Any, _context: Any) -> dict[str, str]: + executor.on_invocation_start("partial-setup", True, datetime.now(UTC), None) + try: + return executor.run_handler(body) + except RuntimeError as error: + assert error is handler_failure + return output + + token = marker.set("incoming") + try: + for _ in range(2): + if outcome == "retry": + with pytest.raises(InvocationError) as caught: + invoke({}, None) + assert caught.value is failure + else: + assert invoke({}, None) is output + assert marker.get() == "incoming" + with pytest.raises(LookupError): + new_binding.get() + finally: + marker.reset(token) + assert healthy_inputs == [("incoming", None)] * 2 + assert cleanup == (names if stage == "start" else ["healthy"]) * 2 + setup = "start" if stage == "start" else "enter" + expected_setup = [ + ("factory" if name == "bad" and stage == "factory" else setup, name) + for name in names + ] + expected_cleanup = ( + [("end", name) for name in names] if stage == "start" else [("exit", "healthy")] + ) + assert events == (expected_setup + expected_cleanup) * 2 + # A token reset in the wrong Context is caught by the SDK, so explicitly + # check diagnostics and successful cleanup rather than relying on raises. + errors = [ + record.exc_info for record in caplog.records if record.exc_info is not None + ] + assert len(errors) == 2 + assert all(str(error[1]) == "partial plugin setup" for error in errors) + + +@pytest.mark.parametrize("shape", ["helper", "property", "dynamic"]) +def test_unopted_legacy_handler_context_is_never_inspected(shape: str) -> None: + from contextlib import contextmanager + from collections.abc import Iterator + from datetime import UTC, datetime + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + calls: list[str] = [] + + @contextmanager + def helper(_info: InvocationStartInfo) -> Iterator[None]: + calls.append("helper") + yield + + def property_getter(_self: Any) -> Any: + calls.append("property") + return helper + + def dynamic_getter(_self: Any, name: str) -> Any: + if name == "handler_context": + calls.append("dynamic") + return "legacy-business-value" + raise AttributeError(name) + + members_by_shape: dict[str, dict[str, Any]] = { + "helper": {"handler_context": staticmethod(helper)}, + "property": {"handler_context": property(property_getter)}, + "dynamic": {"__getattr__": dynamic_getter}, + } + plugin_type = type( + "Legacy", (DurableInstrumentationPlugin,), members_by_shape[shape] + ) + plugin = plugin_type() + if shape == "dynamic": + assert plugin.handler_context == "legacy-business-value" + calls.clear() + executor = PluginExecutor([plugin]) + with executor.run(): + executor.on_invocation_start("legacy-helper", True, datetime.now(UTC), None) + assert executor.run_handler(lambda: list(calls)) == [] + assert calls == [] + + +@pytest.mark.parametrize("marker", [None, 0, 2, True, "1", property(lambda _: 1)]) +def test_handler_scope_requires_literal_class_local_version(marker: Any) -> None: + from contextlib import nullcontext + from datetime import UTC, datetime + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + calls: list[str] = [] + + def helper(_self: Any, _info: InvocationStartInfo) -> Any: + calls.append("scope") + return nullcontext() + + plugin_type = type( + "Legacy", + (DurableInstrumentationPlugin,), + {"__durable_handler_context_api__": marker, "handler_context": helper}, + ) + executor = PluginExecutor([plugin_type()]) + with executor.run(): + executor.on_invocation_start("invalid-marker", True, datetime.now(UTC), None) + assert executor.run_handler(lambda: "ok") == "ok" + assert calls == [] + + +def test_handler_scope_opt_in_is_not_inherited_or_taken_from_instance() -> None: + from contextlib import contextmanager + from collections.abc import Iterator + from datetime import UTC, datetime + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + marker = contextvars.ContextVar("explicit-scope", default="outside") + events: list[str] = [] + + class OptedPlugin(DurableInstrumentationPlugin): + __durable_handler_context_api__ = 1 + + @contextmanager + def handler_context(self, _info: InvocationStartInfo) -> Iterator[None]: + events.append("enter") + token = marker.set("inside") + try: + yield + finally: + marker.reset(token) + events.append("exit") + + class LegacySubclass(OptedPlugin): + pass + + class ExplicitSubclass(OptedPlugin): + __durable_handler_context_api__ = 1 + + for plugin, expected in [ + (OptedPlugin(), "inside"), + (LegacySubclass(), "outside"), + (ExplicitSubclass(), "inside"), + ]: + # Assigning a marker to an instance cannot accidentally enable the hook. + plugin.__durable_handler_context_api__ = 1 + executor = PluginExecutor([plugin]) + with executor.run(): + executor.on_invocation_start("subclass", True, datetime.now(UTC), None) + assert executor.run_handler(marker.get) == expected + assert marker.get() == "outside" + assert events == ["enter", "exit", "enter", "exit"] + + +def test_handler_opt_in_does_not_trigger_legacy_metaclass_descriptors() -> None: + from datetime import UTC, datetime + from aws_durable_execution_sdk_python.plugin import PluginExecutor + + reads: list[str] = [] + + def namespace(_cls: Any) -> Any: + reads.append("metaclass-dict") + raise RuntimeError("legacy namespace") + + def helper(_self: Any, _info: Any) -> Any: + reads.append("legacy-helper") + raise RuntimeError("legacy helper") + + meta = type("LegacyMeta", (type,), {"__dict__": property(namespace)}) + plugin_type = meta( + "Legacy", (DurableInstrumentationPlugin,), {"handler_context": helper} + ) + executor = PluginExecutor([plugin_type()]) + with executor.run(): + executor.on_invocation_start("metaclass", True, datetime.now(UTC), None) + assert executor.run_handler(lambda: "ok") == "ok" + assert reads == [] diff --git a/pyproject.toml b/pyproject.toml index 7b637fea0..d5377965a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -131,8 +131,10 @@ dependencies = [ test = "pytest packages/aws-durable-execution-sdk-python-examples/test {args}" [tool.hatch.envs.test-pypi-otel] +# Do not inherit the default workspace core when validating a published release. +workspace.members = ["packages/aws-durable-execution-sdk-python-otel"] dependencies = [ - "aws-durable-execution-sdk-python>=2.0.0", + "aws-durable-execution-sdk-python==2.0.0", "opentelemetry-sdk>=1.20.0", "opentelemetry-propagator-aws-xray", "pytest",