diff --git a/.github/workflows/opentelemetry-conformance-tests.yml b/.github/workflows/opentelemetry-conformance-tests.yml index 216490be3..4b4f2215a 100644 --- a/.github/workflows/opentelemetry-conformance-tests.yml +++ b/.github/workflows/opentelemetry-conformance-tests.yml @@ -66,7 +66,7 @@ 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@f18bd0b5f28c5c90e288d0fb8bca08a849b51863 with: language: python runs_on: codebuild-github-actions-runner-${{ github.run_id }}-${{ github.run_attempt }} diff --git a/packages/aws-durable-execution-sdk-python-otel/README.md b/packages/aws-durable-execution-sdk-python-otel/README.md index 2c7e1bb1b..ff9400b56 100644 --- a/packages/aws-durable-execution-sdk-python-otel/README.md +++ b/packages/aws-durable-execution-sdk-python-otel/README.md @@ -343,6 +343,34 @@ context onto every emitted log record using these attributes: These attributes are only set when a valid span context is active, so any log formatter or schema must treat the fields as optional. +## Draft chained-invoke propagation + +Both views supply the synchronous `provide_propagation_metadata` hook with +SDK-owned types. It encodes canonical X-Ray `Root`, the calling operation's +`Parent`, and resolved `Sampled=1` or `Sampled=0`, without creating an extra span. +An existing operation's actual span ID is used; before span creation, its stable +ID is derived from the execution ARN and operation ID. Unrelated ambient spans +cannot replace execution ownership. Inactive/mismatched executions contribute +nothing. Tracer, provider and resource ownership remain unchanged. + +The core now calls the collector only for a new invoke START and persists the +contribution in flat `ChainedInvokeOptions.XAmznTraceId`. Pending and terminal +replay do not recollect metadata; an uncommitted START can be retried. Separate +invokes carry separate operation parents. Public invoke tests cover both views, +sampled/unsampled context, parallel branches, preserved tenant/payload and replay. + +This PR stays draft for [#751](https://github.com/aws/aws-durable-execution-sdk-python/issues/751) +until the public generated model and backend support are available. The normal +botocore request tests intentionally expose the missing field; they do not bypass +the serializer or hide the dependency failure. Python has no distributed-map +model/START path; existing map/parallel APIs are CONTEXT operations. + +On older supported cores without the propagation contract, plugin loading and +existing tracing continue, and the optional hook contributes no metadata. The +new capability requires the coordinated core. Dependency floors and provider API +version are unchanged; deployed downstream topology still needs validation after +model/backend publication. + ## Verification After deploying your function with the plugin configured: 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 40a62ca84..62df8e07c 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 @@ -34,7 +34,7 @@ import datetime import logging import threading -from typing import Any, ClassVar +from typing import Any, ClassVar, TYPE_CHECKING from aws_durable_execution_sdk_python.plugin import ( DurableInstrumentationPlugin, @@ -96,6 +96,19 @@ uninstall_log_filter, ) from aws_durable_execution_sdk_python_otel.provider import create_tracer_provider +from aws_durable_execution_sdk_python_otel.propagation import propagation_metadata + + +if TYPE_CHECKING: + from aws_durable_execution_sdk_python.plugin import ( + PropagationInput, + PropagationMetadata, + ) +else: + from aws_durable_execution_sdk_python import plugin as _core_plugin + + PropagationInput = getattr(_core_plugin, "PropagationInput", Any) + PropagationMetadata = getattr(_core_plugin, "PropagationMetadata", Any) logger = logging.getLogger(__name__) @@ -428,6 +441,33 @@ def _with_sampling(self, parent_context: Context) -> Context: # ------------------------------------------------------------------ # Invocation lifecycle # ------------------------------------------------------------------ + def provide_propagation_metadata( + self, + info: PropagationInput, + ) -> PropagationMetadata | None: + """Describe this operation as the downstream parent without starting a span.""" + execution_context = self._execution_trace_context + if ( + not self._tracing_enabled + or execution_context is None + or info.execution_arn != self._execution_arn + or not info.operation_id + ): + return None + operation_span = self._get_span(info.operation_id) + span_context = ( + operation_span.get_span_context() + if operation_span is not None + else self._operation_span_context(info.operation_id) + ) + if ( + span_context is None + or not span_context.is_valid + or span_context.trace_id != execution_context.trace_id + ): + return None + return propagation_metadata(span_context) + def on_invocation_start(self, info: InvocationStartInfo) -> None: logger.debug("Durable invocation started: %s", info) self._registration_accepted = True 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..4ca6a725b 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 @@ -5,7 +5,7 @@ import datetime import logging import threading -from typing import Any, ClassVar +from typing import Any, ClassVar, TYPE_CHECKING from aws_durable_execution_sdk_python.plugin import ( DurableInstrumentationPlugin, @@ -65,6 +65,19 @@ ) from aws_durable_execution_sdk_python_otel.otel_plugin_config import OtelPluginConfig from aws_durable_execution_sdk_python_otel.provider import create_tracer_provider +from aws_durable_execution_sdk_python_otel.propagation import propagation_metadata + + +if TYPE_CHECKING: + from aws_durable_execution_sdk_python.plugin import ( + PropagationInput, + PropagationMetadata, + ) +else: + from aws_durable_execution_sdk_python import plugin as _core_plugin + + PropagationInput = getattr(_core_plugin, "PropagationInput", Any) + PropagationMetadata = getattr(_core_plugin, "PropagationMetadata", Any) logger = logging.getLogger(__name__) @@ -523,6 +536,33 @@ def _end_span( # ------------------------------------------------------------------ # Plugin lifecycle callbacks # ------------------------------------------------------------------ + def provide_propagation_metadata( + self, + info: PropagationInput, + ) -> PropagationMetadata | None: + """Describe this operation as the downstream parent without starting a span.""" + execution_context = self._execution_trace_context + if ( + not self._tracing_enabled + or execution_context is None + or info.execution_arn != self._execution_arn + or not info.operation_id + ): + return None + operation_span = self._get_span(info.operation_id) + span_context = ( + operation_span.get_span_context() + if operation_span is not None + else self._operation_link_context(info.operation_id) + ) + if ( + span_context is None + or not span_context.is_valid + or span_context.trace_id != execution_context.trace_id + ): + return None + return propagation_metadata(span_context) + def on_invocation_start(self, info: InvocationStartInfo) -> None: """Called at the start of each invocation. Creates the invocation span.""" logger.debug("Durable invocation started: %s", info) diff --git a/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/propagation.py b/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/propagation.py new file mode 100644 index 000000000..0cdc3d659 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-otel/src/aws_durable_execution_sdk_python_otel/propagation.py @@ -0,0 +1,31 @@ +"""Pure X-Ray propagation encoding for the SDK-owned plugin contract.""" + +from __future__ import annotations + +from typing import Any, TYPE_CHECKING + +from aws_durable_execution_sdk_python import plugin as core_plugin +from opentelemetry.trace import SpanContext + + +if TYPE_CHECKING: + from aws_durable_execution_sdk_python.plugin import PropagationMetadata +else: + PropagationMetadata = getattr(core_plugin, "PropagationMetadata", Any) + + +def propagation_metadata(span_context: SpanContext) -> PropagationMetadata | None: + """Encode an operation's context without creating a span or changing state.""" + # Old supported cores do not expose the additive contract. Existing tracing + # still works; only this new optional contribution is unavailable there. + metadata_type = getattr(core_plugin, "PropagationMetadata", None) + if metadata_type is None: + return None + trace_id = f"{span_context.trace_id:032x}" + return metadata_type( + x_amzn_trace_id=( + f"Root=1-{trace_id[:8]}-{trace_id[8:]};" + f"Parent={span_context.span_id:016x};" + f"Sampled={int(span_context.trace_flags.sampled)}" + ) + ) diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_invoke_checkpoint_propagation.py b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_invoke_checkpoint_propagation.py new file mode 100644 index 000000000..0f44a7a38 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_invoke_checkpoint_propagation.py @@ -0,0 +1,247 @@ +"""Headers on real public invokes identify each calling operation's actual span.""" + +from __future__ import annotations + +import json +import logging +from datetime import UTC, datetime +from typing import Any +from unittest.mock import Mock, patch + +import pytest +from aws_durable_execution_sdk_python import DurableContext, durable_execution +from aws_durable_execution_sdk_python.config import ( + InvokeConfig, + ParallelBranch, + ParallelConfig, +) +from aws_durable_execution_sdk_python.lambda_service import ( + CheckpointOutput, + CheckpointUpdatedExecutionState, + Operation, + OperationType, + OperationUpdate, +) +from aws_durable_execution_sdk_python.plugin import ( + OperationStartInfo, + PropagationInput, + PropagationMetadata, + OperationType as HookOperationType, +) +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 opentelemetry.trace import SpanContext + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + ExtractedContext, + OtelPluginConfig, + Sampling, +) + + +ARN = ( + "arn:aws:lambda:us-west-2:123456789012:function:parent:1/durable-execution/test/id" +) +TRACE = 0x68E1BE000123456789ABCDEF01234567 + + +@pytest.mark.parametrize("plugin_type", [ExecutionOtelPlugin, InvocationOtelPlugin]) +@pytest.mark.parametrize("sampled", [False, True]) +def test_parallel_invokes_carry_distinct_actual_operation_parents( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + plugin_type: type[ExecutionOtelPlugin] | type[InvocationOtelPlugin], + sampled: bool, +) -> 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, + enrich_logger=False, + context_extractor=lambda _: ExtractedContext( + TRACE, 42, Sampling.SAMPLED if sampled else Sampling.NOT_SAMPLED + ), + ) + ) + inputs: list[PropagationInput] = [] + actual_contexts: dict[str, SpanContext] = {} + producer = plugin.provide_propagation_metadata + start = plugin.on_operation_start + + def produce(info: PropagationInput) -> PropagationMetadata | None: + inputs.append(info) + return producer(info) + + def started(info: OperationStartInfo) -> None: + start(info) + if info.operation_type is HookOperationType.CHAINED_INVOKE: + span = plugin._get_span(info.operation_id) + assert span is not None + actual_contexts[info.operation_id] = span.get_span_context() + + monkeypatch.setattr(plugin, "provide_propagation_metadata", produce) + monkeypatch.setattr(plugin, "on_operation_start", started) + updates_seen: list[OperationUpdate] = [] + operations: dict[str, Operation] = {} + timestamp = datetime.now(UTC) + + def checkpoint( + durable_execution_arn: str, + checkpoint_token: str, + updates: list[OperationUpdate], + client_token: str | None, + ) -> CheckpointOutput: + assert durable_execution_arn == ARN + changed: list[Operation] = [] + for update in updates: + updates_seen.append(update) + assert update.sub_type is not None + raw: dict[str, Any] = { + "Id": update.operation_id, + "Type": update.operation_type.value, + "Name": update.name, + "SubType": update.sub_type.value, + "ParentId": update.parent_id, + "StartTimestamp": timestamp, + "Status": "STARTED" if update.action.value == "START" else "SUCCEEDED", + } + if update.operation_type is OperationType.CHAINED_INVOKE: + assert update.chained_invoke_options is not None + raw.update( + Status="SUCCEEDED", + EndTimestamp=timestamp, + ChainedInvokeDetails={ + "Result": json.dumps( + update.chained_invoke_options.function_name + ) + }, + ) + elif update.action.value == "SUCCEED": + raw.update( + EndTimestamp=timestamp, ContextDetails={"Result": update.payload} + ) + operations[update.operation_id] = Operation.from_dict(raw) + changed.append(operations[update.operation_id]) + return CheckpointOutput( + "next-token", CheckpointUpdatedExecutionState(operations=changed) + ) + + def branch_a(durable: DurableContext) -> str: + return durable.invoke( + "child-a:live", + {"branch": "a"}, + name="invoke-a", + config=InvokeConfig(tenant_id="tenant-a"), + ) + + def branch_b(durable: DurableContext) -> str: + return durable.invoke( + "child-b:live", + {"branch": "b"}, + name="invoke-b", + config=InvokeConfig(tenant_id="tenant-b"), + ) + + def user_handler(_event: Any, durable: DurableContext) -> list[str]: + return durable.parallel( + [ + ParallelBranch(branch_a, "branch-a"), + ParallelBranch(branch_b, "branch-b"), + ], + name="invoke-group", + config=ParallelConfig(max_concurrency=2), + ).get_results() + + handler = durable_execution(user_handler, plugins=[plugin]) + invocation: dict[str, Any] = { + "DurableExecutionArn": ARN, + "CheckpointToken": "checkpoint-token", + "InitialExecutionState": { + "Operations": [ + { + "Id": "id", + "Type": "EXECUTION", + "Status": "STARTED", + "StartTimestamp": int(timestamp.timestamp() * 1000), + "ExecutionDetails": {"InputPayload": "{}"}, + } + ], + "NextMarker": "", + }, + } + lambda_context = Mock() + lambda_context.aws_request_id = "request" + lambda_context.client_context = None + lambda_context.identity = None + lambda_context._epoch_deadline_time_in_ms = 0 + lambda_context.invoked_function_arn = "parent:1" + lambda_context.tenant_id = None + service = Mock() + service.checkpoint = checkpoint + before = context.get_current() + try: + with provider.get_tracer("unrelated").start_as_current_span( + "ambient" + ) as ambient: + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient.initialize_client", + return_value=service, + ): + output = handler(invocation, lambda_context) + assert output["Status"] == "SUCCEEDED", output + assert json.loads(output["Result"]) == ["child-a:live", "child-b:live"] + invocation["InitialExecutionState"]["Operations"].extend( + op.to_json_dict() for op in operations.values() + ) + assert handler(invocation, lambda_context) == output + assert trace.get_current_span() is ambient + assert len(inputs) == 2 + invokes = [ + u for u in updates_seen if u.operation_type is OperationType.CHAINED_INVOKE + ] + assert len(invokes) == 2 and len(actual_contexts) == 2 + assert len({item.parent_operation_id for item in inputs}) == 2 + headers = [] + for update in invokes: + options = update.chained_invoke_options + assert options is not None and options.x_amzn_trace_id is not None + headers.append(options.x_amzn_trace_id) + info = next( + item for item in inputs if item.operation_id == update.operation_id + ) + assert info == PropagationInput( + ARN, update.operation_id, options.function_name, update.parent_id + ) + suffix = "a" if options.function_name == "child-a:live" else "b" + assert options.tenant_id == f"tenant-{suffix}" + assert update.payload is not None + assert json.loads(update.payload) == {"branch": suffix} + fields = dict( + part.split("=", 1) for part in options.x_amzn_trace_id.split(";") + ) + actual = actual_contexts[update.operation_id] + assert actual.trace_id == TRACE + assert fields["Root"] == "1-68e1be00-0123456789abcdef01234567" + assert int(fields["Parent"], 16) == actual.span_id + assert actual.span_id != 42 + assert fields["Sampled"] == str(int(sampled)) + assert len(set(headers)) == 2 + emitted = [ + s + for s in exporter.get_finished_spans() + if s.name in ("invoke-a", "invoke-b") + ] + assert len(emitted) == (2 if sampled else 0) + assert context.get_current() == before + assert not [ + record for record in caplog.records if record.levelno >= logging.ERROR + ] + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_propagation_metadata.py b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_propagation_metadata.py new file mode 100644 index 000000000..0539eed75 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-otel/tests/e2e/test_propagation_metadata.py @@ -0,0 +1,349 @@ +"""Pure propagation production from both views and the actual operation context.""" + +from collections.abc import Iterator +from datetime import UTC, datetime, timedelta + +import pytest +from aws_durable_execution_sdk_python import plugin as core_plugin +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + InvocationStartInfo, + InvocationEndInfo, + InvocationStatus, + OperationStartInfo, + OperationEndInfo, + OperationStatus, + OperationSubType, + OperationType, + PluginExecutor, + PropagationInput, +) +from opentelemetry import context, trace +from opentelemetry.sdk.trace import TracerProvider, SpanProcessor, Span +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.sdk.trace.sampling import ALWAYS_OFF, ALWAYS_ON, Sampler +from opentelemetry.trace import NoOpTracerProvider + +from aws_durable_execution_sdk_python_otel import ( + ExecutionOtelPlugin, + InvocationOtelPlugin, + OtelPluginConfig, + ExtractedContext, + Sampling, +) +from aws_durable_execution_sdk_python_otel.deterministic_id_generator import ( + operation_id_to_span_id, +) + + +START = datetime(2026, 10, 2, tzinfo=UTC) +END = START + timedelta(seconds=1) +ARN = "arn:aws:lambda:us-west-2:123456789012:function:propagation:1/durable-execution/test/id" +TRACE_ID = 0x12345678901234567890123456789012 +INPUT = PropagationInput(ARN, "invoke-1", "target:1") +PluginType = type[ExecutionOtelPlugin] | type[InvocationOtelPlugin] + + +class _Starts(SpanProcessor): + def __init__(self) -> None: + self.count = 0 + + def on_start( + self, span: Span, parent_context: context.Context | None = None + ) -> None: + self.count += 1 + + +@pytest.fixture(params=[ExecutionOtelPlugin, InvocationOtelPlugin]) +def plugin_type(request: pytest.FixtureRequest) -> PluginType: + return request.param + + +@pytest.fixture(autouse=True) +def balanced_context() -> Iterator[None]: + before = context.get_current() + yield + assert context.get_current() == before + + +def start_info(first: bool = True) -> InvocationStartInfo: + return InvocationStartInfo( + request_id="request", + execution_arn=ARN, + execution_start_time=START, + is_first_invocation=first, + ) + + +def end_info( + status: InvocationStatus = InvocationStatus.SUCCEEDED, +) -> InvocationEndInfo: + return InvocationEndInfo( + request_id="request", + execution_arn=ARN, + execution_start_time=START, + is_first_invocation=True, + status=status, + ) + + +def start_operation( + plugin: ExecutionOtelPlugin | InvocationOtelPlugin, replayed: bool = False +) -> None: + plugin.on_operation_start( + OperationStartInfo( + operation_id=INPUT.operation_id, + operation_type=OperationType.CHAINED_INVOKE, + sub_type=OperationSubType.CHAINED_INVOKE, + name="invoke-target", + parent_id=None, + start_time=START, + is_replayed=replayed, + status=OperationStatus.STARTED, + ) + ) + + +def end_operation(plugin: ExecutionOtelPlugin | InvocationOtelPlugin) -> None: + plugin.on_operation_end( + OperationEndInfo( + operation_id=INPUT.operation_id, + operation_type=OperationType.CHAINED_INVOKE, + sub_type=OperationSubType.CHAINED_INVOKE, + name="invoke-target", + parent_id=None, + start_time=START, + end_time=END, + is_replayed=False, + status=OperationStatus.SUCCEEDED, + ) + ) + + +@pytest.mark.parametrize( + ("extracted", "sampler", "sampled"), + [ + (ExtractedContext(TRACE_ID, 42, Sampling.SAMPLED), ALWAYS_OFF, True), + (ExtractedContext(TRACE_ID, 42, Sampling.NOT_SAMPLED), ALWAYS_ON, False), + (None, ALWAYS_ON, True), + (None, ALWAYS_OFF, False), + ], +) +def test_collector_encodes_operation_parent_without_side_effects( + plugin_type: PluginType, + extracted: ExtractedContext | None, + sampler: Sampler, + sampled: bool, +) -> None: + provider = TracerProvider(sampler=sampler) + exporter = InMemorySpanExporter() + starts = _Starts() + provider.add_span_processor(starts) + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: extracted, + enrich_logger=False, + ) + ) + collector = PluginExecutor([DurableInstrumentationPlugin(), plugin]) + assert collector.provide_propagation_metadata(INPUT).x_amzn_trace_id is None + plugin.on_invocation_start(start_info()) + try: + starts_before = starts.count + current_before = trace.get_current_span() + exported_before = exporter.get_finished_spans() + metadata = collector.provide_propagation_metadata(INPUT) + assert starts.count == starts_before + assert trace.get_current_span() is current_before + assert exporter.get_finished_spans() == exported_before + assert metadata.x_amzn_trace_id is not None + fields = dict( + part.split("=", 1) for part in metadata.x_amzn_trace_id.split(";") + ) + assert fields["Sampled"] == ("1" if sampled else "0") + parent_id = int(fields["Parent"], 16) + assert parent_id == operation_id_to_span_id(ARN, INPUT.operation_id) + assert parent_id != 42 + if extracted is not None: + assert fields["Root"] == "1-12345678-901234567890123456789012" + start_operation(plugin) + starts_before = starts.count + # A different active ambient trace must not replace execution ownership. + ambient_provider = TracerProvider() + try: + with ambient_provider.get_tracer("unrelated").start_as_current_span( + "ambient", context=context.Context() + ) as ambient: + assert collector.provide_propagation_metadata(INPUT) == metadata + assert trace.get_current_span() is ambient + finally: + ambient_provider.shutdown() + assert starts.count == starts_before + assert ( + collector.provide_propagation_metadata( + PropagationInput("other", "invoke-1", "target") + ).x_amzn_trace_id + is None + ) + end_operation(plugin) + plugin.on_invocation_end(end_info()) + if sampled: + operation = next( + s for s in exporter.get_finished_spans() if s.name == "invoke-target" + ) + assert operation.context is not None + assert operation.context.span_id == parent_id + trace_hex = f"{operation.context.trace_id:032x}" + assert fields["Root"] == f"1-{trace_hex[:8]}-{trace_hex[8:]}" + else: + assert not exporter.get_finished_spans() + assert collector.provide_propagation_metadata(INPUT).x_amzn_trace_id is None + finally: + plugin.on_invocation_end(end_info()) + provider.shutdown() + + +def test_invocation_continuation_uses_actual_fresh_operation_span() -> None: + provider = TracerProvider() + exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = InvocationOtelPlugin( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + try: + plugin.on_invocation_start(start_info()) + start_operation(plugin) + first = plugin.provide_propagation_metadata(INPUT) + plugin.on_invocation_end(end_info(InvocationStatus.PENDING)) + plugin.on_invocation_start(start_info(first=False)) + start_operation(plugin, replayed=True) + continued = plugin.provide_propagation_metadata(INPUT) + assert first is not None and continued is not None + assert continued.x_amzn_trace_id is not None + assert first.x_amzn_trace_id != continued.x_amzn_trace_id + parent = int( + dict(p.split("=", 1) for p in continued.x_amzn_trace_id.split(";"))[ + "Parent" + ], + 16, + ) + end_operation(plugin) + plugin.on_invocation_end(end_info()) + operations = [ + s for s in exporter.get_finished_spans() if s.name == "invoke-target" + ] + assert operations[-1].context is not None + assert operations[-1].context.span_id == parent + finally: + provider.shutdown() + + +def test_unbound_provider_does_not_produce_metadata( + plugin_type: PluginType, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(trace, "get_tracer_provider", lambda: NoOpTracerProvider()) + plugin = plugin_type( + OtelPluginConfig( + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + plugin.on_invocation_start(start_info()) + assert plugin.provide_propagation_metadata(INPUT) is None + plugin.on_invocation_end(end_info()) + + +def test_resume_preserves_logical_identity_and_rejects_old_execution( + plugin_type: PluginType, +) -> None: + provider = TracerProvider() + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + try: + plugin.on_invocation_start(start_info()) + first = plugin.provide_propagation_metadata(INPUT) + plugin.on_invocation_end(end_info(InvocationStatus.RETRY)) + plugin.on_invocation_start(start_info(first=False)) + assert plugin.provide_propagation_metadata(INPUT) == first + plugin.on_invocation_end(end_info()) + other_start = InvocationStartInfo( + request_id="other", + execution_arn="other-execution", + execution_start_time=START, + is_first_invocation=True, + ) + plugin.on_invocation_start(other_start) + assert plugin.provide_propagation_metadata(INPUT) is None + plugin.on_invocation_end(end_info()) + finally: + provider.shutdown() + + +def test_failed_invocation_setup_cannot_supply_stale_metadata( + plugin_type: PluginType, +) -> None: + provider = TracerProvider() + fail = False + + def extract(_info: InvocationStartInfo) -> ExtractedContext | None: + if fail: + raise ValueError("invalid extraction") + return None + + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, context_extractor=extract, enrich_logger=False + ) + ) + try: + plugin.on_invocation_start(start_info()) + assert plugin.provide_propagation_metadata(INPUT) is not None + plugin.on_invocation_end(end_info()) + fail = True + with pytest.raises(ValueError, match="invalid extraction"): + plugin.on_invocation_start(start_info(first=False)) + assert plugin.provide_propagation_metadata(INPUT) is None + plugin.on_invocation_end(end_info()) + finally: + provider.shutdown() + + +def test_legacy_core_keeps_tracing_without_propagation_contract( + plugin_type: PluginType, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delattr(core_plugin, "PropagationMetadata") + provider = TracerProvider() + exporter = InMemorySpanExporter() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + plugin = plugin_type( + OtelPluginConfig( + tracer_provider=provider, + context_extractor=lambda _: None, + enrich_logger=False, + ) + ) + try: + plugin.on_invocation_start(start_info()) + assert plugin.provide_propagation_metadata(INPUT) is None + start_operation(plugin) + end_operation(plugin) + plugin.on_invocation_end(end_info()) + assert {"invoke-target", "Invocation", "Workflow"} <= { + s.name for s in exporter.get_finished_spans() + } + finally: + provider.shutdown() diff --git a/packages/aws-durable-execution-sdk-python/README.md b/packages/aws-durable-execution-sdk-python/README.md index 18bc0d5e7..88a610a75 100644 --- a/packages/aws-durable-execution-sdk-python/README.md +++ b/packages/aws-durable-execution-sdk-python/README.md @@ -112,6 +112,42 @@ 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. +### Draft chained-invoke propagation + +A new invoke START collects synchronous plugin metadata after resolving its +operation and target identity, before checkpointing. The SDK writes a non-blank +`x_amzn_trace_id` contribution to the flat +`ChainedInvokeOptions.XAmznTraceId` member. Function, tenant, payload, operation +name and identity remain unchanged. The member is omitted without a contribution. +Each invoke in a batch has its own metadata; this is not a request-wide header. + +The SDK-owned frozen `PropagationInput` carries `execution_arn`, `operation_id`, +optional `parent_operation_id`, and `target_function_name`. Frozen +`PropagationMetadata` has optional `x_amzn_trace_id`. Neither type depends on +OpenTelemetry or generated service models. The optional synchronous plugin +`provide_propagation_metadata(info)` hook defaults to no contribution. + +The collector keeps the first non-blank opaque value in configured order without +trimming it. Equal values do not conflict; different later values log both plugin +identities and a conflict count. Ordinary hook/getter/result/diagnostic failures +are isolated; cancellation and other `BaseException` control signals keep their +existing behavior. Replaying a checkpointed START, pending operation or terminal +result does not call the hook. If a START was never committed, a later attempt +can collect again; the callback is not exactly-once. + +This remains draft pending public Lambda model and backend publication. The SDK +model and START path are implemented; the real botocore serialization tests +intentionally fail while `ChainedInvokeOptions.XAmznTraceId` is absent from the +installed model. No field removal, validation bypass or model-capability fallback +is used to make those checks pass. Python has no `DistributedMapOptions` wrapper +or distributed-map START API: its existing map/parallel operations use CONTEXT. +The design's corresponding distributed-map field awaits that future modeled path. + +Release coordination still requires the compatible core/OTel minor versions. +Existing valid core/plugin combinations retain prior tracing behavior; older +cores without this optional contract contribute no new propagation metadata. +Backend rollout and deployed downstream trace-topology validation remain pending. + ## 🚀 Quick Start Install the execution SDK: diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py index eb5ee78b8..cc58dcd65 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py @@ -453,19 +453,18 @@ def to_dict(self) -> MutableMapping[str, Any]: @dataclass(frozen=True) class ChainedInvokeOptions: - """ - As of 2025/10/27: - - Chained invoke options only contains a function name - """ + """Target and optional per-operation trace context for a chained invocation.""" function_name: str tenant_id: str | None = None + x_amzn_trace_id: str | None = None @classmethod def from_dict(cls, data: MutableMapping[str, Any]) -> ChainedInvokeOptions: return cls( function_name=data["FunctionName"], tenant_id=data.get("TenantId"), + x_amzn_trace_id=data.get("XAmznTraceId"), ) def to_dict(self) -> MutableMapping[str, Any]: @@ -474,6 +473,8 @@ def to_dict(self) -> MutableMapping[str, Any]: } if self.tenant_id is not None: result["TenantId"] = self.tenant_id + if self.x_amzn_trace_id is not None: + result["XAmznTraceId"] = self.x_amzn_trace_id return result diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/invoke.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/invoke.py index c64013975..4d8be6d36 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/invoke.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/invoke.py @@ -123,12 +123,18 @@ def check_result_status(self) -> CheckResult[R]: operation_id=self.operation_identifier.operation_id, durable_execution_arn=self.state.durable_execution_arn, ) + # Only a new START contributes outbound context. Replayed pending or + # terminal invokes keep their checkpointed identity and outcome. + propagation = self.state.provide_propagation_metadata( + self.operation_identifier, self.function_name + ) start_operation: OperationUpdate = OperationUpdate.create_invoke_start( identifier=self.operation_identifier, payload=serialized_payload, chained_invoke_options=ChainedInvokeOptions( function_name=self.function_name, tenant_id=self.config.tenant_id, + x_amzn_trace_id=propagation.x_amzn_trace_id, ), ) # Checkpoint invoke START with blocking (is_sync=True). 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..247fb5716 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 @@ -4,6 +4,7 @@ import copy import datetime import functools +import inspect import logging from collections.abc import Mapping, Sequence from concurrent.futures import ThreadPoolExecutor @@ -381,9 +382,41 @@ def from_durable_execution_invocation_output( ) +@dataclass(frozen=True) +class PropagationInput: + """SDK-owned identity for a new chained-invoke START propagation request. + + This contract does not contain generated service-model or telemetry types. + """ + + execution_arn: str + operation_id: str + target_function_name: str + parent_operation_id: str | None = None + + +@dataclass(frozen=True) +class PropagationMetadata: + """Immutable plugin contribution; members are opaque to the core SDK.""" + + x_amzn_trace_id: str | None = None + + class DurableInstrumentationPlugin: """Base class for plugins. Override only the methods you need.""" + def provide_propagation_metadata( + self, + info: PropagationInput, + ) -> PropagationMetadata | None: + """Synchronously provide metadata without performing a durable operation. + + Called before checkpointing a new invoke START, never for its replay. + Existing plugins can inherit the no-op default. A failed, uncommitted + START may collect again on a later attempt. + """ + return None + def on_invocation_start(self, info: InvocationStartInfo) -> None: """Called when an invocation starts. This is called within the thread that runs user function handler. @@ -467,6 +500,83 @@ def __init__(self, plugins: list[DurableInstrumentationPlugin] | None): self._invocation_status: InvocationStartInfo | None = None self._operations_provider: Callable[[], Mapping[str, Operation]] | None = None + def provide_propagation_metadata( + self, info: PropagationInput + ) -> PropagationMetadata: + """Collect opaque metadata synchronously in configured plugin order. + + The first non-blank value wins. Ordinary plugin failures are isolated, + matching lifecycle dispatch; BaseException cancellation/control signals + retain the existing propagation policy. The caller attaches the result + to its operation model before checkpoint serialization. + """ + value: str | None = None + owner: str | None = None + conflict_count = 0 + for plugin in self._plugins: + # Read type identity without invoking instance attribute access. + identity = "unknown plugin" + try: + plugin_type = type(plugin) + identity = type.__getattribute__(plugin_type, "__qualname__") + module = type.__getattribute__(plugin_type, "__module__") + if type(module) is str: + identity = ".".join((module, identity)) + except Exception: + pass + try: + metadata = plugin.provide_propagation_metadata(info) + if metadata is None: + continue + if inspect.iscoroutine(metadata): + metadata.close() + raise TypeError("propagation metadata hook must be synchronous") + if not isinstance(metadata, PropagationMetadata): + raise TypeError("expected PropagationMetadata or None") + candidate = metadata.x_amzn_trace_id + if candidate is not None: + if not isinstance(candidate, str): + raise TypeError("x_amzn_trace_id must be a string or None") + # Treat str subclasses as data, not executable comparison + # hooks during first-value/conflict aggregation. + candidate = str.__str__(candidate) + except Exception: + self._log_propagation_diagnostic( + logging.ERROR, + "Plugin %s propagation metadata failure ignored", + identity, + exc_info=True, + ) + continue + if candidate is None or not candidate.strip(): + continue + if value is None: + value, owner = candidate, identity + elif candidate != value: + conflict_count += 1 + self._log_propagation_diagnostic( + logging.WARNING, + "Propagation metadata conflict between %s and %s for " + "x_amzn_trace_id; keeping first value (conflict_count=%d)", + owner, + identity, + conflict_count, + ) + return PropagationMetadata(x_amzn_trace_id=value) + + @staticmethod + def _log_propagation_diagnostic( + level: int, + message: str, + *args: object, + exc_info: bool = False, + ) -> None: + """Keep ordinary diagnostic failures from changing propagation collection.""" + try: + logger.log(level, message, *args, exc_info=exc_info) + except Exception: + pass + @contextlib.contextmanager def run(self): if self._plugins: diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py index f5ce7214a..85171d37a 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py @@ -38,6 +38,8 @@ ) from aws_durable_execution_sdk_python.plugin import ( PluginExecutor, + PropagationInput, + PropagationMetadata, UserFunctionOutcome, ) from aws_durable_execution_sdk_python.threading import CompletionEvent @@ -589,6 +591,19 @@ def emit_operation_replay_hook(self, operation: Operation) -> None: self._plugin_executor.on_operation_replay(operation) + def provide_propagation_metadata( + self, identifier: OperationIdentifier, target_function_name: str + ) -> PropagationMetadata: + """Collect metadata for a new operation before its START is checkpointed.""" + return self._plugin_executor.provide_propagation_metadata( + PropagationInput( + execution_arn=self.durable_execution_arn, + operation_id=identifier.operation_id, + parent_operation_id=identifier.parent_id, + target_function_name=target_function_name, + ) + ) + def emit_child_context_end_hook( self, operation_identifier: OperationIdentifier, diff --git a/packages/aws-durable-execution-sdk-python/tests/e2e/invoke_propagation_int_test.py b/packages/aws-durable-execution-sdk-python/tests/e2e/invoke_propagation_int_test.py new file mode 100644 index 000000000..c39a3c4ba --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/e2e/invoke_propagation_int_test.py @@ -0,0 +1,272 @@ +"""Public invoke/checkpoint/replay behavior with a recording service boundary. + +These tests exercise the SDK lifecycle, not generated wire support. The latter +is checked independently against the installed client in invoke_wire_propagation_test. +""" + +from __future__ import annotations + +import json +from dataclasses import replace +from datetime import UTC, datetime +from typing import Any +from unittest.mock import Mock, patch + +import pytest + +from aws_durable_execution_sdk_python import DurableContext, durable_execution +from aws_durable_execution_sdk_python.config import InvokeConfig +from aws_durable_execution_sdk_python.exceptions import CheckpointError +from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeDetails, + CheckpointOutput, + CheckpointUpdatedExecutionState, + ErrorObject, + Operation, + OperationStatus, + OperationUpdate, +) +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + PropagationInput, + PropagationMetadata, +) + + +ARN = ( + "arn:aws:lambda:us-west-2:123456789012:function:parent:1/durable-execution/test/id" +) +START = datetime(2026, 10, 5, tzinfo=UTC) + + +def lambda_context() -> Mock: + context = Mock() + context.aws_request_id = "request" + context.client_context = None + context.identity = None + context._epoch_deadline_time_in_ms = 0 + context.invoked_function_arn = ( + "arn:aws:lambda:us-west-2:123456789012:function:parent:1" + ) + context.tenant_id = "runtime-tenant" + return context + + +def event(operations: list[Operation] | None = None) -> dict[str, Any]: + return { + "DurableExecutionArn": ARN, + "CheckpointToken": "checkpoint-token", + "InitialExecutionState": { + "Operations": [ + { + "Id": "id", + "Type": "EXECUTION", + "Status": "STARTED", + "StartTimestamp": int(START.timestamp() * 1000), + "ExecutionDetails": {"InputPayload": "{}"}, + }, + *(operation.to_json_dict() for operation in operations or []), + ], + "NextMarker": "", + }, + } + + +class RecordingService: + def __init__(self) -> None: + self.updates: list[OperationUpdate] = [] + self.operations: list[Operation] = [] + self.fail = False + + def checkpoint( + self, + durable_execution_arn: str, + checkpoint_token: str, + updates: list[OperationUpdate], + client_token: str | None, + ) -> CheckpointOutput: + assert durable_execution_arn == ARN + assert checkpoint_token + self.updates.extend(updates) + if self.fail: + raise CheckpointError("uncommitted START") + for update in updates: + self.operations.append( + Operation( + operation_id=update.operation_id, + operation_type=update.operation_type, + parent_id=update.parent_id, + name=update.name, + sub_type=update.sub_type, + status=OperationStatus.STARTED, + start_timestamp=START, + ) + ) + return CheckpointOutput( + "next-token", + CheckpointUpdatedExecutionState(operations=self.operations.copy()), + ) + + +class RecordingPlugin(DurableInstrumentationPlugin): + def __init__(self, mode: str = "value") -> None: + self.calls: list[PropagationInput] = [] + self.mode = mode + + def provide_propagation_metadata( + self, info: PropagationInput + ) -> PropagationMetadata | None: + self.calls.append(info) + if self.mode == "error": + raise ValueError("instrumentation failed") + if self.mode == "absent": + return None + if self.mode == "blank": + return PropagationMetadata(" \t\n") + return PropagationMetadata(f"header-{len(self.calls)}") + + +def body(_event: Any, context: DurableContext) -> Any: + return context.invoke( + "child:live", + {"payload": [1, 2]}, + name="call-child", + config=InvokeConfig(tenant_id="explicit-tenant"), + ) + + +@pytest.mark.parametrize( + "terminal", + [ + OperationStatus.SUCCEEDED, + OperationStatus.FAILED, + OperationStatus.TIMED_OUT, + OperationStatus.STOPPED, + ], +) +def test_public_invoke_pending_and_terminal_replay_do_not_recollect( + monkeypatch: pytest.MonkeyPatch, terminal: OperationStatus +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + plugin = RecordingPlugin() + service = RecordingService() + handler = durable_execution(body, plugins=[plugin]) + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient.initialize_client", + return_value=service, + ): + first = handler(event(), lambda_context()) + assert first["Status"] == "PENDING" + assert len(plugin.calls) == 1 and len(service.updates) == 1 + update = service.updates[0] + assert plugin.calls == [ + PropagationInput(ARN, update.operation_id, "child:live", None) + ] + assert update.name == "call-child" + assert update.payload is not None + assert json.loads(update.payload) == {"payload": [1, 2]} + assert update.to_dict()["ChainedInvokeOptions"] == { + "FunctionName": "child:live", + "TenantId": "explicit-tenant", + "XAmznTraceId": "header-1", + } + pending = handler(event(service.operations), lambda_context()) + assert pending["Status"] == "PENDING" + saved = replace( + service.operations[0], + status=terminal, + end_timestamp=START, + chained_invoke_details=ChainedInvokeDetails( + result='"saved child result"' + if terminal is OperationStatus.SUCCEEDED + else None, + error=None + if terminal is OperationStatus.SUCCEEDED + else ErrorObject("saved child error", "ChildError", None, None), + ), + ) + result = handler(event([saved]), lambda_context()) + assert len(plugin.calls) == 1 and len(service.updates) == 1 + if terminal is OperationStatus.SUCCEEDED: + assert result == {"Status": "SUCCEEDED", "Result": '"saved child result"'} + else: + assert result["Status"] == "FAILED" + assert result["Error"]["ErrorMessage"] == "saved child error" + assert result["Error"]["ErrorType"].endswith("InvokeError") + + +@pytest.mark.parametrize( + "mode", ["no-plugins", "absent-hook", "absent", "blank", "error"] +) +def test_public_invoke_without_a_contribution_keeps_legacy_request( + monkeypatch: pytest.MonkeyPatch, mode: str +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + plugin = RecordingPlugin(mode) + plugins: list[DurableInstrumentationPlugin] = ( + [] + if mode == "no-plugins" + else [DurableInstrumentationPlugin()] + if mode == "absent-hook" + else [plugin] + ) + service = RecordingService() + handler = durable_execution(body, plugins=plugins) + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient.initialize_client", + return_value=service, + ): + assert handler(event(), lambda_context())["Status"] == "PENDING" + assert len(service.updates) == 1 + assert service.updates[0].to_dict()["ChainedInvokeOptions"] == { + "FunctionName": "child:live", + "TenantId": "explicit-tenant", + } + + +def test_public_invoke_recollects_after_failed_uncommitted_checkpoint( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + plugin = RecordingPlugin() + service = RecordingService() + handler = durable_execution(body, plugins=[plugin]) + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient.initialize_client", + return_value=service, + ): + service.fail = True + with pytest.raises(CheckpointError, match="uncommitted START"): + handler(event(), lambda_context()) + service.fail = False + assert handler(event(), lambda_context())["Status"] == "PENDING" + assert ( + handler(event(service.operations), lambda_context())["Status"] == "PENDING" + ) + assert len(plugin.calls) == 2 + assert plugin.calls[0] == plugin.calls[1] + headers: list[str | None] = [] + for update in service.updates: + assert update.chained_invoke_options is not None + headers.append(update.chained_invoke_options.x_amzn_trace_id) + assert headers == ["header-1", "header-2"] + assert service.updates[0].operation_id == service.updates[1].operation_id + + +def test_blank_plugin_does_not_mask_later_header_on_real_start( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("DURABLE_EXECUTION_PLUGINS", raising=False) + blank, healthy = RecordingPlugin("blank"), RecordingPlugin() + service = RecordingService() + handler = durable_execution(body, plugins=[blank, healthy]) + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient.initialize_client", + return_value=service, + ): + assert handler(event(), lambda_context())["Status"] == "PENDING" + assert blank.calls == healthy.calls + assert ( + service.updates[0].to_dict()["ChainedInvokeOptions"]["XAmznTraceId"] + == "header-1" + ) diff --git a/packages/aws-durable-execution-sdk-python/tests/invoke_wire_propagation_test.py b/packages/aws-durable-execution-sdk-python/tests/invoke_wire_propagation_test.py new file mode 100644 index 000000000..0b6491e4e --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/invoke_wire_propagation_test.py @@ -0,0 +1,138 @@ +"""Owned-model round trips and the real installed Lambda wire serializer. + +The header-bearing client tests intentionally fail until botocore publishes +ChainedInvokeOptions.XAmznTraceId. They are not skipped or given preview models. +""" + +from __future__ import annotations + +import json +from io import BytesIO +from collections.abc import Iterator +from unittest.mock import patch +from typing import Any + +import boto3 +import pytest +from botocore.awsrequest import AWSResponse +from botocore.config import Config +from botocore.compat import HTTPHeaders + +from aws_durable_execution_sdk_python.identifier import OperationIdentifier +from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeOptions, + LambdaClient, + OperationSubType, + OperationUpdate, +) + + +ARN = ( + "arn:aws:lambda:us-west-2:123456789012:function:parent:1/durable-execution/test/id" +) +HEADER = "Root=1-68e1be00-0123456789abcdef01234567;Parent=0123456789abcdef;Sampled=1" + + +@pytest.mark.parametrize("header", [None, "", HEADER]) +@pytest.mark.parametrize("tenant", [None, "tenant-a"]) +def test_chained_options_and_update_round_trip( + header: str | None, tenant: str | None +) -> None: + options = ChainedInvokeOptions("child:live", tenant, header) + expected = {"FunctionName": "child:live"} + if tenant is not None: + expected["TenantId"] = tenant + if header is not None: + expected["XAmznTraceId"] = header + assert options.to_dict() == expected + assert ChainedInvokeOptions.from_dict(expected) == options + update = OperationUpdate.create_invoke_start( + OperationIdentifier( + "invoke-a", OperationSubType.CHAINED_INVOKE, "group", "child" + ), + '{"preserved": true}', + options, + ) + assert OperationUpdate.from_dict(update.to_dict()) == update + assert update.to_dict()["ChainedInvokeOptions"] == expected + assert update.to_dict()["Payload"] == '{"preserved": true}' + + +class _ResponseBody(BytesIO): + def stream(self, amt: int = 1024, decode_content: bool = False) -> Iterator[bytes]: + yield self.read() + + +@pytest.mark.parametrize("parameter_validation", [True, False]) +@pytest.mark.parametrize("with_header", [False, True]) +def test_public_botocore_serializes_each_invokes_header( + parameter_validation: bool, with_header: bool +) -> None: + """Exercise LambdaClient -> ordinary boto client -> HTTP request, without AWS I/O.""" + client = boto3.client( + "lambda", + region_name="us-west-2", + aws_access_key_id="testing", + aws_secret_access_key="testing", + endpoint_url="https://lambda.example.invalid", + config=Config( + parameter_validation=parameter_validation, retries={"max_attempts": 0} + ), + ) + requests: list[dict[str, Any]] = [] + + def capture(request: Any, **_kwargs: Any) -> AWSResponse: + requests.append(json.loads(request.body)) + response_headers = HTTPHeaders() + response_headers["content-type"] = "application/json" + return AWSResponse( + request.url, + 200, + response_headers, + _ResponseBody( + b'{"CheckpointToken":"next","NewExecutionState":{"Operations":[]}}' + ), + ) + + headers = [ + HEADER, + HEADER.replace("0123456789abcdef;Sampled", "fedcba9876543210;Sampled"), + ] + updates = [ + OperationUpdate.create_invoke_start( + OperationIdentifier( + f"invoke-{index}", + OperationSubType.CHAINED_INVOKE, + "group", + f"child-{index}", + ), + json.dumps({"index": index}), + ChainedInvokeOptions( + f"child-{index}:live", "tenant-a", header if with_header else None + ), + ) + for index, header in enumerate(headers) + ] + try: + # Replace only network I/O; parameter validation and request serialization + # are the unmodified installed botocore path. + with patch("botocore.httpsession.URLLib3Session.send", side_effect=capture): + output = LambdaClient(client).checkpoint( + ARN, "checkpoint", updates, "client-token" + ) + assert output.checkpoint_token == "next" + assert len(requests) == 1 + assert requests[0]["CheckpointToken"] == "checkpoint" + assert requests[0]["ClientToken"] == "client-token" + assert requests[0]["Updates"] == [update.to_dict() for update in updates] + for index, update in enumerate(requests[0]["Updates"]): + options = update["ChainedInvokeOptions"] + assert options["FunctionName"] == f"child-{index}:live" + assert options["TenantId"] == "tenant-a" + assert json.loads(update["Payload"]) == {"index": index} + if with_header: + assert options["XAmznTraceId"] == headers[index] + else: + assert "XAmznTraceId" not in options + finally: + client.close() diff --git a/packages/aws-durable-execution-sdk-python/tests/operation/invoke_propagation_test.py b/packages/aws-durable-execution-sdk-python/tests/operation/invoke_propagation_test.py new file mode 100644 index 000000000..c0c35ca64 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/operation/invoke_propagation_test.py @@ -0,0 +1,169 @@ +"""A propagation hook belongs to a new invoke START, never to a replay read.""" + +from __future__ import annotations + +import asyncio +from unittest.mock import Mock + +import pytest + +from aws_durable_execution_sdk_python.config import InvokeConfig +from aws_durable_execution_sdk_python.exceptions import ( + CheckpointError, + InvokeError, + SuspendExecution, +) +from aws_durable_execution_sdk_python.identifier import OperationIdentifier +from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeDetails, + ErrorObject, + Operation, + OperationStatus, + OperationSubType, + OperationType, +) +from aws_durable_execution_sdk_python.operation.invoke import InvokeOperationExecutor +from aws_durable_execution_sdk_python.plugin import PropagationMetadata +from aws_durable_execution_sdk_python.state import CheckpointedResult, ExecutionState + + +IDENTIFIER = OperationIdentifier( + "invoke", OperationSubType.CHAINED_INVOKE, "parent", "named-invoke" +) + + +def _executor(state: Mock) -> InvokeOperationExecutor: + return InvokeOperationExecutor( + "child:live", + {"payload": [1, 2]}, + state, + IDENTIFIER, + InvokeConfig(tenant_id="tenant-a"), + ) + + +def test_new_start_collects_before_checkpoint_and_preserves_owned_fields() -> None: + state = Mock(spec=ExecutionState) + state.durable_execution_arn = "execution" + state.get_checkpoint_result.return_value = CheckpointedResult.create_not_found() + events: list[str] = [] + + def collect(*args): + events.append("collect") + assert args == (IDENTIFIER, "child:live") + return PropagationMetadata("opaque-header") + + state.provide_propagation_metadata.side_effect = collect + state.create_checkpoint.side_effect = lambda **_kwargs: events.append("checkpoint") + assert not _executor(state).check_result_status().is_ready_to_execute + assert events == ["collect", "checkpoint"] + update = state.create_checkpoint.call_args.kwargs["operation_update"] + assert state.create_checkpoint.call_args.kwargs["is_sync"] is True + assert update.operation_id == "invoke" and update.parent_id == "parent" + assert update.name == "named-invoke" + assert update.to_dict()["Payload"] == '{"payload": [1, 2]}' + assert update.to_dict()["ChainedInvokeOptions"] == { + "FunctionName": "child:live", + "TenantId": "tenant-a", + "XAmznTraceId": "opaque-header", + } + + +@pytest.mark.parametrize( + "status", + [ + OperationStatus.STARTED, + OperationStatus.PENDING, + OperationStatus.SUCCEEDED, + OperationStatus.FAILED, + OperationStatus.TIMED_OUT, + OperationStatus.STOPPED, + ], +) +def test_checkpointed_start_or_terminal_replay_never_collects( + status: OperationStatus, +) -> None: + state = Mock(spec=ExecutionState) + state.durable_execution_arn = "execution" + state.get_checkpoint_result.return_value = CheckpointedResult.create_from_operation( + Operation( + operation_id="invoke", + operation_type=OperationType.CHAINED_INVOKE, + sub_type=OperationSubType.CHAINED_INVOKE, + parent_id="parent", + name="named-invoke", + status=status, + chained_invoke_details=ChainedInvokeDetails( + result='"saved"', + error=ErrorObject( + message="saved failure", + type="SavedError", + data=None, + stack_trace=None, + ), + ), + ) + ) + if status is OperationStatus.SUCCEEDED: + assert _executor(state).process() == "saved" + elif status in (OperationStatus.STARTED, OperationStatus.PENDING): + with pytest.raises(SuspendExecution): + _executor(state).process() + else: + with pytest.raises(InvokeError, match="saved failure"): + _executor(state).process() + state.provide_propagation_metadata.assert_not_called() + state.create_checkpoint.assert_not_called() + + +def test_uncommitted_start_can_recollect_after_checkpoint_failure() -> None: + state = Mock(spec=ExecutionState) + state.durable_execution_arn = "execution" + state.get_checkpoint_result.return_value = CheckpointedResult.create_not_found() + state.provide_propagation_metadata.side_effect = [ + PropagationMetadata("first"), + PropagationMetadata("retry"), + ] + state.create_checkpoint.side_effect = CheckpointError("not committed") + for _ in range(2): + with pytest.raises(CheckpointError, match="not committed"): + _executor(state).check_result_status() + updates = [ + call.kwargs["operation_update"] + for call in state.create_checkpoint.call_args_list + ] + assert [update.chained_invoke_options.x_amzn_trace_id for update in updates] == [ + "first", + "retry", + ] + assert all(update.operation_id == "invoke" for update in updates) + + +@pytest.mark.parametrize( + "failure", [KeyboardInterrupt, SystemExit, GeneratorExit, asyncio.CancelledError] +) +def test_fatal_hook_failure_prevents_start_checkpoint( + failure: type[BaseException], +) -> None: + from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + PluginExecutor, + PropagationInput, + ) + + class FatalPlugin(DurableInstrumentationPlugin): + def provide_propagation_metadata( + self, info: PropagationInput + ) -> PropagationMetadata: + raise failure("control flow") + + service = Mock() + state = ExecutionState( + "execution", "token", {}, service, PluginExecutor([FatalPlugin()]) + ) + executor: InvokeOperationExecutor[str] = InvokeOperationExecutor( + "child:live", {}, state, IDENTIFIER, InvokeConfig() + ) + with pytest.raises(failure, match="control flow"): + executor.check_result_status() + service.checkpoint.assert_not_called() diff --git a/packages/aws-durable-execution-sdk-python/tests/operation/invoke_test.py b/packages/aws-durable-execution-sdk-python/tests/operation/invoke_test.py index 6fca520b3..e2ad3cd91 100644 --- a/packages/aws-durable-execution-sdk-python/tests/operation/invoke_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/operation/invoke_test.py @@ -24,10 +24,17 @@ OperationSubType, ) from aws_durable_execution_sdk_python.operation.invoke import InvokeOperationExecutor +from aws_durable_execution_sdk_python.plugin import PropagationMetadata from aws_durable_execution_sdk_python.state import CheckpointedResult, ExecutionState from tests.serdes_test import CustomDictSerDes +def _state_without_plugins() -> Mock: + state = Mock(spec=ExecutionState) + state.provide_propagation_metadata.return_value = PropagationMetadata() + return state + + # Test helper - maintains old handler signature for backward compatibility in tests def invoke_handler(function_name, payload, state, operation_identifier, config): """Test helper that wraps InvokeOperationExecutor with old handler signature.""" @@ -45,7 +52,7 @@ def invoke_handler(function_name, payload, state, operation_identifier, config): def test_invoke_handler_already_succeeded(): """Test invoke_handler when operation already succeeded.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -75,7 +82,7 @@ def test_invoke_handler_already_succeeded(): def test_invoke_handler_already_succeeded_none_result(): """Test invoke_handler when operation succeeded with None result.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -104,7 +111,7 @@ def test_invoke_handler_already_succeeded_none_result(): def test_invoke_handler_already_succeeded_no_chained_invoke_details(): """Test invoke_handler when operation succeeded but has no chained_invoke_details.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -136,7 +143,7 @@ def test_invoke_handler_already_succeeded_no_chained_invoke_details(): ) def test_invoke_handler_already_terminated(kind: OperationStatus): """Test invoke_handler when operation already failed.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" error = ErrorObject( @@ -167,7 +174,7 @@ def test_invoke_handler_already_terminated(kind: OperationStatus): def test_invoke_handler_already_timed_out(): """Test invoke_handler when operation already timed out.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" error = ErrorObject( @@ -199,7 +206,7 @@ def test_invoke_handler_already_timed_out(): @pytest.mark.parametrize("status", [OperationStatus.STARTED]) def test_invoke_handler_already_started(status): """Test invoke_handler when operation is already started.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -230,7 +237,7 @@ def test_invoke_handler_already_started(status): @pytest.mark.parametrize("status", [OperationStatus.STARTED, OperationStatus.PENDING]) def test_invoke_handler_already_started_suspends(status): """Test invoke_handler when operation is already started suspends indefinitely.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -261,7 +268,7 @@ def test_invoke_handler_already_started_suspends(status): def test_invoke_handler_new_operation(): """Test invoke_handler when starting a new operation.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: started (no immediate response) @@ -305,7 +312,7 @@ def test_invoke_handler_new_operation(): def test_invoke_handler_no_config(): """Test invoke_handler when no config is provided.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -340,7 +347,7 @@ def test_invoke_handler_no_config(): def test_invoke_handler_custom_serdes(): """Test invoke_handler with custom serialization.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -376,7 +383,7 @@ def test_invoke_handler_custom_serdes(): def test_invoke_handler_custom_serdes_new_operation(): """Test invoke_handler with custom serialization for new operation.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -415,7 +422,7 @@ def test_invoke_handler_custom_serdes_new_operation(): @pytest.mark.parametrize("status", [OperationStatus.STARTED, OperationStatus.PENDING]) def test_invoke_handler_with_operation_name(status: OperationStatus): """Test invoke_handler uses operation name in logs when available.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -444,7 +451,7 @@ def test_invoke_handler_with_operation_name(status: OperationStatus): @pytest.mark.parametrize("status", [OperationStatus.STARTED, OperationStatus.PENDING]) def test_invoke_handler_without_operation_name(status: OperationStatus): """Test invoke_handler uses function name in logs when no operation name.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -472,7 +479,7 @@ def test_invoke_handler_without_operation_name(status: OperationStatus): def test_invoke_handler_with_none_payload(): """Test invoke_handler when payload is None.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -505,7 +512,7 @@ def test_invoke_handler_with_none_payload(): def test_invoke_handler_already_succeeded_with_none_payload(): """Test invoke_handler when operation succeeded and original payload was None.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" operation = Operation( @@ -539,7 +546,7 @@ def test_invoke_handler_already_succeeded_with_none_payload(): def test_invoke_handler_suspend_does_not_raise(mock_suspend): """Test invoke_handler when suspend_with_optional_resume_delay doesn't raise an exception.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -575,7 +582,7 @@ def test_invoke_handler_suspend_does_not_raise(mock_suspend): def test_invoke_handler_with_tenant_id(): """Test invoke_handler passes tenant_id to checkpoint.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -612,7 +619,7 @@ def test_invoke_handler_with_tenant_id(): def test_invoke_handler_without_tenant_id(): """Test invoke_handler without tenant_id doesn't include it in checkpoint.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -649,7 +656,7 @@ def test_invoke_handler_without_tenant_id(): def test_invoke_handler_default_config_no_tenant_id(): """Test invoke_handler with default config has no tenant_id.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -684,7 +691,7 @@ def test_invoke_handler_default_config_no_tenant_id(): def test_invoke_handler_defaults_to_json_serdes(): """Test invoke_handler uses DEFAULT_JSON_SERDES when config has no serdes.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" not_found = CheckpointedResult.create_not_found() @@ -719,7 +726,7 @@ def test_invoke_handler_defaults_to_json_serdes(): def test_invoke_handler_result_defaults_to_json_serdes(): """Test invoke_handler uses DEFAULT_JSON_SERDES for result deserialization.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" result_data = {"key": "value", "number": 42} @@ -757,7 +764,7 @@ def test_invoke_handler_result_defaults_to_json_serdes(): def test_invoke_immediate_response_get_checkpoint_result_called_twice(): """Test that get_checkpoint_result is called twice when checkpoint is created.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: started (no immediate response) @@ -792,7 +799,7 @@ def test_invoke_immediate_response_get_checkpoint_result_called_twice(): def test_invoke_immediate_response_create_checkpoint_with_is_sync_true(): """Test that create_checkpoint is called with is_sync=True.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: started @@ -833,7 +840,7 @@ def test_invoke_immediate_response_immediate_success(): When checkpoint returns SUCCEEDED on second check, operation returns result without suspend. """ - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: succeeded (immediate response) @@ -871,7 +878,7 @@ def test_invoke_immediate_response_immediate_success(): def test_invoke_immediate_response_immediate_success_with_none_result(): """Test immediate success with None result.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: succeeded with None result @@ -912,7 +919,7 @@ def test_invoke_immediate_response_immediate_failure(status: OperationStatus): When checkpoint returns a failure status on second check, operation raises error without suspend. """ - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: failed (immediate response) @@ -957,7 +964,7 @@ def test_invoke_immediate_response_no_immediate_response(): When checkpoint returns STARTED on second check, operation suspends normally. """ - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: started (no immediate response) @@ -999,7 +1006,7 @@ def test_invoke_immediate_response_already_completed(): When checkpoint is already SUCCEEDED on first check, no checkpoint is created and result is returned immediately. """ - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: already succeeded @@ -1036,7 +1043,7 @@ def test_invoke_immediate_response_already_completed(): def test_invoke_immediate_response_with_custom_serdes(): """Test immediate success with custom serialization.""" - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: not found, second call: succeeded @@ -1079,7 +1086,7 @@ def test_invoke_suspends_when_second_check_returns_started(): Validates: Requirements 8.1, 8.2 """ - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: checkpoint doesn't exist @@ -1119,7 +1126,7 @@ def test_invoke_suspends_when_second_check_returns_started_duplicate(): """Test backward compatibility: when the second checkpoint check returns STARTED (not terminal), the invoke operation suspends normally. """ - mock_state = Mock(spec=ExecutionState) + mock_state = _state_without_plugins() mock_state.durable_execution_arn = "test_arn" # First call: checkpoint doesn't exist diff --git a/packages/aws-durable-execution-sdk-python/tests/propagation_test.py b/packages/aws-durable-execution-sdk-python/tests/propagation_test.py new file mode 100644 index 000000000..1fe9eaac3 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/propagation_test.py @@ -0,0 +1,213 @@ +"""Synchronous propagation collection without generated client dependencies.""" + +import asyncio +import logging +import threading +from dataclasses import FrozenInstanceError +from typing import Any +from unittest.mock import patch + +import pytest + +from aws_durable_execution_sdk_python.plugin import ( + DurableInstrumentationPlugin, + PluginExecutor, + PropagationInput, + PropagationMetadata, +) + + +INFO = PropagationInput("execution", "invoke-1", "target:1", "parent-1") + + +class _Alpha(DurableInstrumentationPlugin): + def __init__(self, value: str | None) -> None: + self.value = value + + def provide_propagation_metadata( + self, info: PropagationInput + ) -> PropagationMetadata: + return PropagationMetadata(self.value) + + +class _Beta(_Alpha): + pass + + +class _Gamma(_Alpha): + pass + + +class _Broken(DurableInstrumentationPlugin): + def provide_propagation_metadata( + self, info: PropagationInput + ) -> PropagationMetadata: + raise ValueError("broken hook") + + +class _BrokenHookGetter(DurableInstrumentationPlugin): + def __getattribute__(self, name: str) -> Any: + if name == "provide_propagation_metadata": + raise ValueError("broken hook lookup") + return super().__getattribute__(name) + + +class _BrokenMetadata(PropagationMetadata): + def __getattribute__(self, name: str) -> Any: + if name == "x_amzn_trace_id": + raise ValueError("broken result getter") + return super().__getattribute__(name) + + +class _InvalidResult(DurableInstrumentationPlugin): + def __init__(self, result: Any) -> None: + self.result = result + + def provide_propagation_metadata(self, info: PropagationInput) -> Any: + return self.result + + +def test_contract_is_immutable_and_existing_plugins_default_to_no_metadata() -> None: + with pytest.raises(FrozenInstanceError): + INFO.operation_id = "changed" # type: ignore[misc] + metadata = PropagationMetadata("header") + with pytest.raises(FrozenInstanceError): + metadata.x_amzn_trace_id = "changed" # type: ignore[misc] + legacy = DurableInstrumentationPlugin() + assert legacy.provide_propagation_metadata(INFO) is None + assert ( + PluginExecutor([legacy]).provide_propagation_metadata(INFO) + == PropagationMetadata() + ) + + +def test_first_non_blank_wins_and_equal_values_do_not_conflict( + caplog: pytest.LogCaptureFixture, +) -> None: + result = PluginExecutor( + [_Alpha(None), _Beta("first"), _Gamma("first")] + ).provide_propagation_metadata(INFO) + assert result == PropagationMetadata("first") + assert "conflict" not in caplog.text + assert PluginExecutor([_Alpha(""), _Beta("later")]).provide_propagation_metadata( + INFO + ) == PropagationMetadata("later") + + +def test_conflicts_name_both_plugins_and_count( + caplog: pytest.LogCaptureFixture, +) -> None: + result = PluginExecutor( + [_Alpha("first"), _Beta("second"), _Gamma("third")] + ).provide_propagation_metadata(INFO) + assert result == PropagationMetadata("first") + warnings = [r.message for r in caplog.records if r.levelno == logging.WARNING] + assert len(warnings) == 2 + assert "_Alpha" in warnings[0] and "_Beta" in warnings[0] + assert "conflict_count=1" in warnings[0] + assert "_Alpha" in warnings[1] and "_Gamma" in warnings[1] + assert "conflict_count=2" in warnings[1] + + +@pytest.mark.parametrize( + "broken", + [ + _Broken(), + _BrokenHookGetter(), + _InvalidResult({"x_amzn_trace_id": "unsupported"}), + _InvalidResult(PropagationMetadata(123)), # type: ignore[arg-type] + _InvalidResult(_BrokenMetadata("header")), + ], +) +def test_ordinary_failures_leave_healthy_plugins_runnable( + broken: DurableInstrumentationPlugin, +) -> None: + result = PluginExecutor([broken, _Alpha("healthy")]).provide_propagation_metadata( + INFO + ) + assert result == PropagationMetadata("healthy") + + +def test_collection_is_synchronous_and_uses_readonly_input() -> None: + calls: list[tuple[int, PropagationInput]] = [] + + class Capturing(DurableInstrumentationPlugin): + def provide_propagation_metadata( + self, info: PropagationInput + ) -> PropagationMetadata: + calls.append((threading.get_ident(), info)) + return PropagationMetadata("header") + + assert PluginExecutor([Capturing()]).provide_propagation_metadata( + INFO + ) == PropagationMetadata("header") + assert calls == [(threading.get_ident(), INFO)] + + +def test_diagnostic_failures_do_not_replace_healthy_results() -> None: + with patch( + "aws_durable_execution_sdk_python.plugin.logger.log", + side_effect=ValueError("broken log handler"), + ): + result = PluginExecutor( + [_Broken(), _Alpha("first"), _Beta("second")] + ).provide_propagation_metadata(INFO) + assert result == PropagationMetadata("first") + + +def test_accidental_async_return_is_rejected_and_closed() -> None: + async def metadata() -> PropagationMetadata: + return PropagationMetadata("async") + + coroutine = metadata() + result = PluginExecutor( + [_InvalidResult(coroutine), _Alpha("sync")] + ).provide_propagation_metadata(INFO) + assert result == PropagationMetadata("sync") + assert getattr(coroutine, "cr_frame") is None + + +@pytest.mark.parametrize( + "failure", [KeyboardInterrupt, SystemExit, GeneratorExit, asyncio.CancelledError] +) +def test_existing_fatal_and_cancellation_policy_is_preserved( + failure: type[BaseException], +) -> None: + class Fatal(DurableInstrumentationPlugin): + def provide_propagation_metadata( + self, info: PropagationInput + ) -> PropagationMetadata: + raise failure() + + with pytest.raises(failure): + PluginExecutor([Fatal(), _Alpha("later")]).provide_propagation_metadata(INFO) + + +def test_string_subclass_comparison_cannot_break_aggregation() -> None: + class StringWithHooks(str): + def __ne__(self, other: object) -> bool: + raise ValueError("unexpected comparison hook") + + def __str__(self) -> str: + return self + + result = PluginExecutor( + [ + _Alpha(StringWithHooks("first")), + _Beta(StringWithHooks("second")), + ] + ).provide_propagation_metadata(INFO) + assert result == PropagationMetadata("first") + assert type(result.x_amzn_trace_id) is str + + +@pytest.mark.parametrize("empty", ["", " ", "\t\n"]) +def test_blank_contribution_does_not_block_later_opaque_value(empty: str) -> None: + value = " opaque-header " + assert PluginExecutor([_Alpha(empty), _Beta(value)]).provide_propagation_metadata( + INFO + ) == PropagationMetadata(value) + assert ( + PluginExecutor([_Alpha(empty)]).provide_propagation_metadata(INFO) + == PropagationMetadata() + )