Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -350,7 +350,7 @@ def on_invocation_end(self, info: InvocationEndInfo) -> None:
invocation_span.set_attribute(
"durable.invocation.status", info.status.value
)
if info.status is InvocationStatus.FAILED:
if info.status in (InvocationStatus.FAILED, InvocationStatus.RETRY):
invocation_span.set_status(
StatusCode.ERROR, info.error.message if info.error else ""
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -175,7 +175,7 @@ def test_invocation_span_records_subsequent_invocation():
("invocation_status", "expected_span_status"),
[
(InvocationStatus.PENDING, StatusCode.UNSET),
(InvocationStatus.RETRY, StatusCode.UNSET),
(InvocationStatus.RETRY, StatusCode.ERROR),

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think we decided to map RETRY to UNSET.

(InvocationStatus.SUCCEEDED, StatusCode.OK),
(InvocationStatus.FAILED, StatusCode.ERROR),
],
Expand All @@ -184,7 +184,7 @@ def test_invocation_span_status_reflects_execution_status(
invocation_status: InvocationStatus,
expected_span_status: StatusCode,
):
"""Only terminal invocation spans receive a success or failure status."""
"""Invocation spans reflect successful, failed, and retry outcomes."""
plugin, exporter = _create_plugin()

plugin.on_invocation_start(_invocation_start_info())
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
BackgroundThreadError,
BotoClientError,
CheckpointError,
DurableExecutionsError,
ExecutionError,
InvocationError,
SuspendExecution,
Expand Down Expand Up @@ -424,6 +425,11 @@ def wrapper(event: Any, context: LambdaContext) -> MutableMapping[str, Any]:
# all user-space errors go here
logger.exception("Execution failed")

if not isinstance(e, DurableExecutionsError):
# Preserve the terminal execution response while identifying
# the uncaught handler error in invocation telemetry.
plugin_executor.override_invocation_status(InvocationStatus.RETRY)

result = DurableExecutionInvocationOutput(
status=InvocationStatus.FAILED, error=ErrorObject.from_exception(e)
).to_dict()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -265,6 +265,7 @@ def __init__(self, plugins: list[DurableInstrumentationPlugin] | None):
self._plugins = plugins or []
self._executor: ThreadPoolExecutor | None = None
self._invocation_status: InvocationStartInfo | None = None
self._invocation_status_override: InvocationStatus | None = None

@contextlib.contextmanager
def run(self):
Expand All @@ -277,6 +278,7 @@ def run(self):
yield
finally:
self._invocation_status = None
self._invocation_status_override = None
# Shut down the thread pool, waiting for pending tasks to complete.
if self._executor:
self._executor.shutdown(wait=True)
Expand Down Expand Up @@ -331,8 +333,13 @@ def on_invocation_start(
is_first_invocation=is_first_invocation,
execution_start_time=execution_start_time,
)
self._invocation_status_override = None
self.execute_plugins(self._invocation_status, sync=True)

def override_invocation_status(self, status: InvocationStatus) -> None:
"""Override the status reported to plugins without changing the output."""
self._invocation_status_override = status

def on_invocation_end(
self,
output: "DurableExecutionInvocationOutput",
Expand All @@ -341,6 +348,14 @@ def on_invocation_end(
# on_invocation_start not called, skip
return

if self._invocation_status_override is not None:
output = DurableExecutionInvocationOutput(
status=self._invocation_status_override,
result=output.result,
error=output.error,
)
self._invocation_status_override = None

invocation_end_info = (
InvocationEndInfo.from_durable_execution_invocation_output(
self._invocation_status, output
Expand Down
30 changes: 28 additions & 2 deletions packages/aws-durable-execution-sdk-python/tests/execution_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
ExecutionError,
GetExecutionStateError,
InvocationError,
StepError,
SuspendExecution,
)
from aws_durable_execution_sdk_python.execution import (
Expand Down Expand Up @@ -2934,8 +2935,8 @@ def test_handler(event: Any, context: DurableContext) -> dict:
assert "invocation_end:SUCCEEDED" in plugin.calls


def test_durable_execution_with_plugins_failure():
"""Test that plugins receive invocation end and execution end on user error."""
def test_durable_execution_with_plugins_uncaught_error():
"""Test that plugins classify an uncaught user error as a retry."""
mock_client = Mock(spec=DurableServiceClient)
mock_output = CheckpointOutput(
checkpoint_token="new_token", # noqa: S106
Expand All @@ -2955,6 +2956,31 @@ def test_handler(event: Any, context: DurableContext) -> dict:
_make_lambda_context(),
)

assert result["Status"] == InvocationStatus.FAILED.value
assert "invocation_start" in plugin.calls
assert "invocation_end:RETRY" in plugin.calls


def test_durable_execution_with_plugins_operation_failure():
"""Test that plugins classify a durable operation error as failed."""
mock_client = Mock(spec=DurableServiceClient)
mock_output = CheckpointOutput(
checkpoint_token="new_token", # noqa: S106
new_execution_state=CheckpointUpdatedExecutionState(),
)
mock_client.checkpoint.return_value = mock_output

plugin = _RecordingPlugin()

@durable_execution(plugins=[plugin])
def test_handler(event: Any, context: DurableContext) -> dict:
raise StepError("step failed")

result = test_handler(
_make_invocation_input(mock_client),
_make_lambda_context(),
)

assert result["Status"] == InvocationStatus.FAILED.value
assert "invocation_start" in plugin.calls
assert "invocation_end:FAILED" in plugin.calls
Expand Down
Loading