From 69ca3cacaa5a5afe1ec491e86bab5e8cc8d2ea2d Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Wed, 30 Sep 2026 23:26:34 +0000 Subject: [PATCH 01/41] Decode canonical cooperative cancellation identity and ranges --- .../_cooperative_cancellation.py | 217 ++++++++++++++++++ tests/test_cooperative_cancellation.py | 164 +++++++++++++ 2 files changed, 381 insertions(+) create mode 100644 src/durable_workflow/_cooperative_cancellation.py create mode 100644 tests/test_cooperative_cancellation.py diff --git a/src/durable_workflow/_cooperative_cancellation.py b/src/durable_workflow/_cooperative_cancellation.py new file mode 100644 index 0000000..64582bd --- /dev/null +++ b/src/durable_workflow/_cooperative_cancellation.py @@ -0,0 +1,217 @@ +"""Canonical cancellation state used by service workflow replay. + +An observation supplies request identity. Only the durable delivery event +authorizes an exception at an authored workflow call. +""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime +from typing import Any + +from .errors import NonDeterministicReplayError + +REQUEST_EVENT = "CooperativeCancellationRequested" +DELIVERY_EVENT = "CooperativeCancellationDelivered" +_CALL_KINDS = {"activity", "local_activity", "timer", "condition", "signal", "child", "parallel", "selection_handle"} +_RESOLUTION_EVENTS = { + "ActivityCompleted", + "ActivityFailed", + "ActivityCancelled", + "ActivityTimedOut", + "TimerFired", + "TimerCancelled", + "ConditionWaitSatisfied", + "ConditionWaitTimedOut", + "SignalApplied", + "ChildRunCompleted", + "ChildRunFailed", + "ChildRunCancelled", + "ChildRunTerminated", +} + + +def _invalid(detail: str, sequence: int = 0) -> NonDeterministicReplayError: + return NonDeterministicReplayError( + sequence, + "canonical cooperative cancellation", + [REQUEST_EVENT, DELIVERY_EVENT], + detail=detail, + ) + + +def _text(value: Any, field: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise _invalid(f"{field} must be a non-empty string") + return value + + +def _positive(value: Any, field: str, *, maximum: int = 2**63 - 1) -> int: + if type(value) is not int or not 1 <= value <= maximum: + raise _invalid(f"{field} must be a positive integer within {maximum}") + return value + + +def _timestamp(value: Any, field: str) -> str: + text = _text(value, field) + try: + parsed = datetime.fromisoformat(text.replace("Z", "+00:00")) + except ValueError as exc: + raise _invalid(f"{field} must be an ISO timestamp") from exc + if parsed.tzinfo is None: + raise _invalid(f"{field} must include its timezone") + return text + + +@dataclass(frozen=True) +class CancellationRequest: + request_id: str + requested_at: str + cleanup_deadline_at: str + history_refresh_page_token: str | None = None + + @classmethod + def from_observation(cls, value: Mapping[str, Any]) -> CancellationRequest: + request = cls( + _text(value.get("request_id"), "request_id"), + _timestamp(value.get("requested_at"), "requested_at"), + _timestamp(value.get("cleanup_deadline_at"), "cleanup_deadline_at"), + _text(value.get("history_refresh_page_token"), "history_refresh_page_token"), + ) + request._validate_deadline() + return request + + def _validate_deadline(self) -> None: + requested = datetime.fromisoformat(self.requested_at.replace("Z", "+00:00")) + deadline = datetime.fromisoformat(self.cleanup_deadline_at.replace("Z", "+00:00")) + if deadline <= requested: + raise _invalid("cleanup deadline must follow the original request") + + +@dataclass(frozen=True) +class CancellationDelivery: + request_id: str + sequence: int + call_kind: str + sequence_span: int = 1 + operation_sequence: int | None = None + operation_sequence_span: int = 1 + + @classmethod + def from_payload(cls, payload: Mapping[str, Any]) -> CancellationDelivery: + request_id = _text(payload.get("workflow_command_id"), "workflow_command_id") + sequence = _positive(payload.get("sequence"), "sequence") + kind = _text(payload.get("call_kind"), "call_kind") + span = _positive(payload.get("sequence_span", 1), "sequence_span", maximum=1000) + operation = payload.get("operation_sequence") + operation_span = _positive(payload.get("operation_sequence_span", 1), "operation_sequence_span", maximum=1000) + if kind not in _CALL_KINDS or (kind != "parallel" and span != 1): + raise _invalid("delivery kind and span do not describe a durable call", sequence) + if sequence > 2**63 - 1 - span: + raise _invalid("delivery sequence range overflows", sequence) + if kind == "selection_handle": + operation = _positive(operation, "operation_sequence") + if operation >= sequence or operation_span > sequence - operation: + raise _invalid("selection handle must name an earlier authored operation", sequence) + elif operation is not None or operation_span != 1: + raise _invalid("only selection handles may carry an operation range", sequence) + return cls(request_id, sequence, kind, span, operation, operation_span) + + def interrupts(self, sequence: int) -> bool: + base = self.operation_sequence if self.operation_sequence is not None else self.sequence + span = self.operation_sequence_span if self.operation_sequence is not None else self.sequence_span + return base <= sequence < base + span + + +@dataclass(frozen=True) +class CancellationHistory: + request: CancellationRequest | None + delivery: CancellationDelivery | None + request_index: int + delivery_index: int | None + resolved_before_request: frozenset[int] + failed_before_request: frozenset[int] + + def eligible(self, sequence: int, span: int = 1) -> bool: + if self.request is None or self.delivery is not None: + return False + sequences = set(range(sequence, sequence + span)) + return not sequences <= self.resolved_before_request and not sequences & self.failed_before_request + + +def read_cancellation_history( + events: Sequence[dict[str, Any]], + *, + run_id: str = "", + observation: Mapping[str, Any] | None = None, +) -> CancellationHistory: + request = CancellationRequest.from_observation(observation) if observation is not None else None + request_index = len(events) + delivery: CancellationDelivery | None = None + delivery_index: int | None = None + saw_request = False + for index, event in enumerate(events): + kind = event.get("event_type") or event.get("type") + if kind not in {REQUEST_EVENT, DELIVERY_EVENT}: + continue + payload = event.get("payload") + if not isinstance(payload, Mapping): + raise _invalid("canonical event is missing its payload") + request_id = _text(payload.get("workflow_command_id"), "workflow_command_id") + if event.get("workflow_command_id", request_id) != request_id: + raise _invalid("event and payload disagree on the original request identity") + event_run = _text(payload.get("workflow_run_id"), "workflow_run_id") + if run_id and event_run != run_id: + raise _invalid("cancellation belongs to a different workflow run") + if kind == REQUEST_EVENT: + if saw_request or delivery is not None: + raise _invalid("history must contain one request before delivery") + canonical_request = CancellationRequest( + request_id, + request.requested_at + if request is not None + else _timestamp( + event.get("recorded_at", event.get("timestamp")), + "requested_at", + ), + _timestamp(payload.get("cleanup_deadline_at"), "cleanup_deadline_at"), + request.history_refresh_page_token if request is not None else None, + ) + if request is not None and ( + request.request_id != canonical_request.request_id + or datetime.fromisoformat(request.cleanup_deadline_at.replace("Z", "+00:00")) + != datetime.fromisoformat(canonical_request.cleanup_deadline_at.replace("Z", "+00:00")) + ): + raise _invalid("observation changes the original request or cleanup deadline") + request = canonical_request + request_index = index + saw_request = True + else: + if not saw_request or request is None or delivery is not None: + raise _invalid("delivery requires one earlier canonical request and one marker") + delivery = CancellationDelivery.from_payload(payload) + if delivery.request_id != request.request_id: + raise _invalid("delivery names a different request", delivery.sequence) + delivery_index = index + resolved: set[int] = set() + failed: set[int] = set() + for event in events[:request_index]: + kind = event.get("event_type") or event.get("type") + payload = event.get("payload") or {} + if not isinstance(payload, Mapping): + continue + sequence = payload.get("sequence") + if type(sequence) is int and sequence > 0 and kind in _RESOLUTION_EVENTS: + resolved.add(sequence) + if kind in { + "ActivityFailed", + "ActivityTimedOut", + "ActivityCancelled", + "ChildRunFailed", + "ChildRunCancelled", + "ChildRunTerminated", + }: + failed.add(sequence) + return CancellationHistory(request, delivery, request_index, delivery_index, frozenset(resolved), frozenset(failed)) diff --git a/tests/test_cooperative_cancellation.py b/tests/test_cooperative_cancellation.py new file mode 100644 index 0000000..6d28bc0 --- /dev/null +++ b/tests/test_cooperative_cancellation.py @@ -0,0 +1,164 @@ +from __future__ import annotations + +from copy import deepcopy +from typing import Any + +import pytest + +from durable_workflow._cooperative_cancellation import CancellationDelivery, read_cancellation_history +from durable_workflow.errors import NonDeterministicReplayError + + +def observation() -> dict[str, Any]: + return { + "request_id": "request-1", + "requested_at": "2026-09-30T12:00:00Z", + "cleanup_deadline_at": "2026-09-30T12:10:00Z", + "history_refresh_page_token": "opaque-first-page", + } + + +def request() -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationRequested", + "workflow_command_id": "request-1", + "recorded_at": "2026-09-30T12:00:00.000120Z", + "payload": { + "workflow_command_id": "request-1", + "workflow_run_id": "run-1", + "cleanup_deadline_at": "2026-09-30T12:10:00Z", + }, + } + + +def marker(sequence: int = 2, call_kind: str = "activity", **fields: Any) -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationDelivered", + "workflow_command_id": "request-1", + "payload": { + "workflow_command_id": "request-1", + "workflow_run_id": "run-1", + "sequence": sequence, + "call_kind": call_kind, + **fields, + }, + } + + +def test_observation_retains_original_identity_without_becoming_delivery() -> None: + state = read_cancellation_history([request()], run_id="run-1", observation=observation()) + assert state.request is not None + assert state.request.requested_at == observation()["requested_at"] + assert state.request.history_refresh_page_token == "opaque-first-page" + assert state.delivery is None + assert state.eligible(1) + + +def test_prior_result_remains_resolved_but_later_result_is_eligible() -> None: + result = {"event_type": "ActivityCompleted", "payload": {"sequence": 1}} + later = {"event_type": "TimerFired", "payload": {"sequence": 2}} + state = read_cancellation_history([result, request(), later]) + assert not state.eligible(1) + assert state.eligible(2) + assert state.eligible(1, 2) + + +def test_prior_parallel_failure_wins_over_a_later_request() -> None: + state = read_cancellation_history( + [ + {"event_type": "ActivityFailed", "payload": {"sequence": 1}}, + request(), + ] + ) + assert not state.eligible(1, 2) + + +def test_cold_history_preserves_one_canonical_delivery_and_its_range() -> None: + state = read_cancellation_history([request(), marker(2, "parallel", sequence_span=3)], run_id="run-1") + assert state.delivery == CancellationDelivery("request-1", 2, "parallel", 3) + assert [state.delivery.interrupts(sequence) for sequence in range(1, 6)] == [False, True, True, True, False] + assert not state.eligible(5) + + +def test_selection_handle_targets_its_earlier_operation_range() -> None: + state = read_cancellation_history( + [ + request(), + marker(5, "selection_handle", operation_sequence=2, operation_sequence_span=2), + ] + ) + assert state.delivery is not None + assert state.delivery.interrupts(2) + assert state.delivery.interrupts(3) + assert not state.delivery.interrupts(5) + + +@pytest.mark.parametrize( + "fields", + [ + {"sequence": True}, + {"sequence": 0}, + {"sequence": -1}, + {"sequence": 2**63 - 1}, + {"call_kind": "side_effect"}, + {"sequence_span": 2}, + {"sequence_span": 0}, + {"call_kind": "parallel", "sequence_span": 1001}, + {"operation_sequence": 1}, + {"operation_sequence_span": 2}, + {"call_kind": "selection_handle"}, + {"call_kind": "selection_handle", "operation_sequence": 2}, + {"call_kind": "selection_handle", "operation_sequence": 1, "operation_sequence_span": 2}, + ], +) +def test_invalid_marker_ranges_fail_closed(fields: dict[str, Any]) -> None: + event = marker() + event["payload"].update(fields) + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([request(), event]) + + +@pytest.mark.parametrize( + "history", + [ + [marker()], + [marker(), request()], + [request(), request()], + [request(), marker(), marker()], + ], +) +def test_missing_reordered_or_duplicate_canonical_markers_fail_closed(history: list[dict[str, Any]]) -> None: + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history(history) + + +@pytest.mark.parametrize("field", ["request_id", "cleanup_deadline_at"]) +def test_observation_cannot_replace_original_request(field: str) -> None: + observed = observation() + observed[field] = "different" if field == "request_id" else "2026-09-30T12:20:00Z" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([request()], observation=observed) + + +@pytest.mark.parametrize( + "field,value", + [ + ("request_id", ""), + ("requested_at", "2026-09-30T12:00:00"), + ("cleanup_deadline_at", "2026-09-30T11:00:00Z"), + ("history_refresh_page_token", ""), + ], +) +def test_invalid_observation_fails_closed(field: str, value: Any) -> None: + observed = observation() + observed[field] = value + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([], observation=observed) + + +def test_delivery_run_and_request_identity_are_checked() -> None: + for field, value in [("workflow_run_id", "other-run"), ("workflow_command_id", "other-request")]: + event = deepcopy(marker()) + event["payload"][field] = value + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([request(), event], run_id="run-1") From 18ec381ea7c0d9378775df033893b9069bdb7802 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Wed, 30 Sep 2026 23:39:59 +0000 Subject: [PATCH 02/41] Replay canonical cancellation at durable calls and recover cleanup --- .../_cooperative_cancellation.py | 29 +- src/durable_workflow/errors.py | 3 +- src/durable_workflow/workflow.py | 273 +++++++++++++++++- ...ooperative-reopened-condition-cleanup.json | 33 +++ tests/test_cooperative_cancellation.py | 263 ++++++++++++++++- tests/test_replay_regression_corpus.py | 15 +- 6 files changed, 596 insertions(+), 20 deletions(-) create mode 100644 tests/fixtures/replay_regressions/cooperative-reopened-condition-cleanup.json diff --git a/src/durable_workflow/_cooperative_cancellation.py b/src/durable_workflow/_cooperative_cancellation.py index 64582bd..34adbbf 100644 --- a/src/durable_workflow/_cooperative_cancellation.py +++ b/src/durable_workflow/_cooperative_cancellation.py @@ -133,12 +133,17 @@ class CancellationHistory: delivery_index: int | None resolved_before_request: frozenset[int] failed_before_request: frozenset[int] + selected_before_request: frozenset[tuple[int, int]] def eligible(self, sequence: int, span: int = 1) -> bool: if self.request is None or self.delivery is not None: return False sequences = set(range(sequence, sequence + span)) - return not sequences <= self.resolved_before_request and not sequences & self.failed_before_request + return ( + not sequences <= self.resolved_before_request + and not sequences & self.failed_before_request + and (sequence, span) not in self.selected_before_request + ) def read_cancellation_history( @@ -197,12 +202,23 @@ def read_cancellation_history( delivery_index = index resolved: set[int] = set() failed: set[int] = set() + selected: set[tuple[int, int]] = set() for event in events[:request_index]: kind = event.get("event_type") or event.get("type") payload = event.get("payload") or {} if not isinstance(payload, Mapping): continue sequence = payload.get("sequence") + if kind == "SelectionResolved": + base = payload.get("selection_group_base_sequence") + span = payload.get("selection_group_size") + if type(base) is int and base > 0 and type(span) is int and 1 <= span <= 1000: + selected.add((base, span)) + elif kind == "SelectionOperationCancelled": + base = payload.get("member_base_sequence") + span = payload.get("member_size") + if type(base) is int and base > 0 and type(span) is int and 1 <= span <= 1000: + resolved.update(range(base, base + span)) if type(sequence) is int and sequence > 0 and kind in _RESOLUTION_EVENTS: resolved.add(sequence) if kind in { @@ -214,4 +230,13 @@ def read_cancellation_history( "ChildRunTerminated", }: failed.add(sequence) - return CancellationHistory(request, delivery, request_index, delivery_index, frozenset(resolved), frozenset(failed)) + state = CancellationHistory( + request, delivery, request_index, delivery_index, frozenset(resolved), frozenset(failed), frozenset(selected), + ) + if delivery is not None: + base = delivery.operation_sequence or delivery.sequence + span = delivery.operation_sequence_span if delivery.operation_sequence is not None else delivery.sequence_span + interrupted = set(range(base, base + span)) + if interrupted <= resolved or interrupted & failed or (base, span) in selected: + raise _invalid("delivery cannot replace an earlier committed result", delivery.sequence) + return state diff --git a/src/durable_workflow/errors.py b/src/durable_workflow/errors.py index 0f8ac20..143a2db 100644 --- a/src/durable_workflow/errors.py +++ b/src/durable_workflow/errors.py @@ -655,8 +655,9 @@ class WorkflowCancelled(BaseException): class by name. """ - def __init__(self, message: str = "workflow was cancelled") -> None: + def __init__(self, message: str = "workflow was cancelled", *, request_id: str | None = None) -> None: super().__init__(message) + self.request_id = request_id class ActivityCancelled(BaseException): diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 4810f43..aa6941d 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -33,6 +33,7 @@ from typing import Any, TypeVar, cast from . import serializer +from ._cooperative_cancellation import CancellationDelivery, read_cancellation_history from .client import WorkflowStreamAppendItem from .errors import ( ActivityFailed, @@ -1680,16 +1681,17 @@ def compensate(self, initiating_failure: BaseException) -> Any: if self._closed: raise RuntimeError("a saga instance can only compensate once") self._closed = True - for compensation in reversed(self._compensations): - try: - yield compensation.command - except (Exception, WorkflowCancelled) as compensation_failure: - raise SagaCompensationFailed( - initiating_failure, - compensation_failure, - compensation_activity_type=compensation.command.activity_type, - compensation_registration_order=compensation.registration_order, - ) from compensation_failure + with self._context.cancellation_shield(): + for compensation in reversed(self._compensations): + try: + yield compensation.command + except (Exception, WorkflowCancelled) as compensation_failure: + raise SagaCompensationFailed( + initiating_failure, + compensation_failure, + compensation_activity_type=compensation.command.activity_type, + compensation_registration_order=compensation.registration_order, + ) from compensation_failure def run(self, forward: Callable[[Saga], Any]) -> Any: """Run ``forward`` and compensate if it fails or is cancelled.""" @@ -1727,6 +1729,8 @@ def __init__( self._current_time = current_time or datetime.now(timezone.utc) self._workflow_command_id = workflow_command_id or run_id or workflow_id self._cancel_requested = bool(cancel_requested) + self._cancellation_request_id: str | None = None + self._cancellation_shield_depth = 0 seed = int(hashlib.sha256(run_id.encode()).hexdigest()[:16], 16) self._rng = random.Random(seed) self._uuid7_counter = 0 @@ -1810,17 +1814,30 @@ def _accept_message_stream(self, arguments: list[Any]) -> None: @property def is_cancellation_requested(self) -> bool: - """Whether this task carries a cooperative cancellation request. + """Whether cancellation has reached this replay's authored boundary. - Server's current ``/cancel`` route is terminal and does not set this - flag. Service-mode cooperative cancellation is not yet available. + Transport observation alone does not set this flag. Server's terminal + ``/cancel`` route does not deliver a cooperative workflow exception. """ return self._cancel_requested def throw_if_cancellation_requested(self) -> None: """Raise :class:`WorkflowCancelled` at an explicit safe point.""" - if self._cancel_requested: - raise WorkflowCancelled("workflow cancellation was requested") + if self._cancel_requested and self._cancellation_shield_depth == 0: + raise WorkflowCancelled("workflow cancellation was requested", request_id=self._cancellation_request_id) + + @contextlib.contextmanager + def cancellation_shield(self) -> Generator[None, None, None]: + """Permit cleanup calls without delivering the same request again. + + The Server still enforces the original cleanup deadline, task lease + and termination fence. This scope does not extend any of them. + """ + self._cancellation_shield_depth += 1 + try: + yield + finally: + self._cancellation_shield_depth -= 1 def saga(self) -> Saga: """Create a deterministic reverse-order compensation scope.""" @@ -2241,6 +2258,7 @@ class ReplayOutcome: commands: list[Command] message_stream_cursors: list[dict[str, Any]] = field(default_factory=list) message_stream_waits: list[dict[str, Any]] = field(default_factory=list) + cancellation_delivery: CancellationDelivery | None = None class Replayer: @@ -2328,6 +2346,7 @@ class _PendingReceiver: name: str args: list[Any] condition_wait_id: str | None = None + after_cancellation_delivery: bool = False @dataclass(frozen=True) @@ -2611,6 +2630,7 @@ def replay( external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, cancel_requested: bool = False, + cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, ) -> ReplayOutcome: return _replay_state( @@ -2624,6 +2644,7 @@ def replay( external_storage=external_storage, external_storage_cache=external_storage_cache, cancel_requested=cancel_requested, + cancellation_request=cancellation_request, local_activity_executor=local_activity_executor, ).outcome @@ -3606,6 +3627,7 @@ def _replay_state( external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, cancel_requested: bool = False, + cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, stop_at_uncommitted_cancellation: bool = False, ) -> _ReplayState: @@ -3622,6 +3644,11 @@ def _replay_state( ) from exception events = list(history_events) + cancellation = read_cancellation_history(events, run_id=run_id, observation=cancellation_request) + cancellation_consumed = False + cancellation_intent: CancellationDelivery | None = None + authored_sequence = 1 + current_call_sequence = 1 workflow_id = workflow_id or _workflow_id_from_history(events) event_types_by_sequence: dict[int, list[str]] = {} details_by_sequence: dict[int, dict[str, Any]] = {} @@ -3725,6 +3752,7 @@ def _state(commands: list[Command]) -> _ReplayState: return _ReplayState( outcome=ReplayOutcome( commands=commands, + cancellation_delivery=cancellation_intent, message_stream_cursors=[ {"stream_name": name, "through_position": position} for name, position in sorted(ctx._message_stream_cursors.items()) @@ -3924,6 +3952,15 @@ def _assert_next_step_matches(command: Any, offset: int = 0) -> None: step_index = result_cursor + offset if step_index >= len(recorded_steps): return + if ( + cancellation.request is not None + and recorded_steps[step_index].workflow_sequence != current_call_sequence + offset + ): + raise NonDeterministicReplayError( + current_call_sequence + offset, _command_diagnostic_shape(command), + recorded_steps[step_index].event_types, + detail="recorded result belongs to a different authored call", + ) _assert_step_matches(command, recorded_steps[step_index]) def _unconsumed_recorded_steps() -> list[_RecordedStep]: @@ -3949,6 +3986,11 @@ def _assert_pending_step_matches(command: Any, offset: int = 0) -> None: _assert_step_matches(command, candidates[offset]) def _assert_no_unconsumed_history(terminal_shape: str) -> None: + if cancellation.delivery is not None and not cancellation_consumed: + raise NonDeterministicReplayError( + cancellation.delivery.sequence, terminal_shape, ["CooperativeCancellationDelivered"], + detail="workflow did not reach the committed cancellation boundary", + ) step = _next_unconsumed_recorded_step() if step is not None: raise NonDeterministicReplayError( @@ -4007,6 +4049,8 @@ def _is_external_receiver_event(event_type: str | None) -> bool: return event_type in ("SignalReceived", "UpdateApplied") def _receiver_binding_boundary_kind(event_type: str | None, payload: Mapping[str, Any]) -> str | None: + if event_type == "CooperativeCancellationDelivered": + return "condition" if payload.get("call_kind") == "condition" else "step" if event_type in ( "ConditionWaitSatisfied", "ConditionWaitTimedOut", @@ -4154,6 +4198,25 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: raw_payload = ev.get("payload") or {} payload = dict(raw_payload) if isinstance(raw_payload, Mapping) else {} sequence = _workflow_sequence(payload) + if ( + cancellation.delivery is not None and sequence is not None + and cancellation.delivery.interrupts(sequence) + and etype in { + "ActivityScheduled", "ActivityStarted", "ActivityCompleted", "ActivityFailed", "ActivityTimedOut", + "ActivityCancelled", "TimerScheduled", "TimerFired", "TimerCancelled", "ChildWorkflowScheduled", + "ChildRunStarted", "ChildRunCompleted", "ChildRunFailed", "ChildRunCancelled", "ChildRunTerminated", + "ConditionWaitOpened", "ConditionWaitSatisfied", "ConditionWaitTimedOut", "SignalWaitOpened", + "SignalApplied", + } + and not (cancellation.delivery.call_kind == "selection_handle" and etype in { + "ActivityScheduled", "ActivityStarted", "TimerScheduled", "ChildWorkflowScheduled", + "ChildRunStarted", "ConditionWaitOpened", "SignalWaitOpened", + }) + and not (cancellation.delivery.call_kind == "condition" and etype == "ConditionWaitOpened") + ): + # The delivery marker owns the interruption. Its activity/timer + # terminal rows are not ordinary workflow results or cleanup calls. + continue if ( etype in { @@ -4365,6 +4428,9 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: external_storage_cache=external_storage_cache, ), condition_wait_id=condition_wait_id, + after_cancellation_delivery=( + cancellation.delivery_index is not None and event_index > cancellation.delivery_index + ), )) elif etype == "UpdateApplied": update_name = payload.get("update_name") @@ -4393,6 +4459,9 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: external_storage_cache=external_storage_cache, ), condition_wait_id=condition_wait_id, + after_cancellation_delivery=( + cancellation.delivery_index is not None and event_index > cancellation.delivery_index + ), )) resolved_results, recorded_steps = _reorder_completed_parallel_groups( @@ -4428,6 +4497,8 @@ def _apply_receiver(receiver: _PendingReceiver) -> None: handler(*receiver.args) def _receiver_due(receiver: _PendingReceiver, *, before_consuming_result: bool) -> bool: + if receiver.after_cancellation_delivery and not cancellation_consumed: + return False if receiver.condition_wait_id is not None: return False if before_consuming_result: @@ -4608,6 +4679,8 @@ def _annotate_parallel_commands( candidates = _unconsumed_recorded_steps() if base_sequence_override is not None: base_sequence = base_sequence_override + elif cancellation.request is not None: + base_sequence = current_call_sequence elif candidates: base_sequence = candidates[0].workflow_sequence else: @@ -4911,6 +4984,8 @@ def _annotate_selection( "SelectionResolved with a positive group base sequence", ["SelectionResolved"], ) + elif cancellation.request is not None: + base_sequence = current_call_sequence elif candidates: base_sequence = candidates[0].workflow_sequence elif any(sequence not in consumed_selection_sequences for sequence in selection_steps): @@ -5034,6 +5109,86 @@ def _consume_terminal_condition_reopens() -> None: break wait_yield_count += 1 + def _cancellation_boundary(command: Any) -> CancellationDelivery | None: + nonlocal authored_sequence, current_call_sequence + if cancellation.request is None: + return None + current_call_sequence = authored_sequence + kind: str | None = None + span = 1 + operation_sequence: int | None = None + operation_span = 1 + if isinstance(command, ScheduleActivity): + kind = "activity" + elif isinstance(command, RecordLocalActivity): + kind = "local_activity" + elif isinstance(command, StartTimer): + kind = "timer" + elif isinstance(command, StartChildWorkflow): + kind = "child" + elif isinstance(command, WaitCondition): + has_recorded_open = any( + _history_event_type(event) == "ConditionWaitOpened" + and _workflow_sequence(event.get("payload") or {}) == current_call_sequence + for event in events + ) + if has_recorded_open or ( + cancellation.delivery is not None and cancellation.delivery.sequence == current_call_sequence + ) or not command.predicate(): + kind = "condition" + elif isinstance(command, list): + kind = "parallel" + span = len(_parallel_leaves(command)) + elif isinstance(command, SelectGroup): + kind = "parallel" + span = sum(len(_parallel_leaves(operation)) if isinstance(operation, list) else 1 + for _, operation in command.operations) + elif isinstance(command, DurableOperationHandle): + kind = "selection_handle" + operation_sequence = command.base_sequence + operation_span = command.size + if kind != "selection_handle" and (kind is not None or isinstance( + command, RecordSideEffect | RecordVersionMarker | UpsertMemo | UpsertSearchAttributes | NexusServiceCall, + )): + authored_sequence += span + if kind is None or span == 0 or cancellation.request is None: + return None + return CancellationDelivery( + cancellation.request.request_id, current_call_sequence, kind, span, + operation_sequence, operation_span, + ) + + def _assert_cancellation_call_matches(command: Any, boundary: CancellationDelivery) -> None: + if isinstance(command, list): + leaves, _ = _annotate_parallel_commands(command, boundary.sequence) + elif isinstance(command, SelectGroup): + leaves, _, _ = _annotate_selection(command, None) + elif isinstance(command, DurableOperationHandle): + _assert_authored_selection_handle(command) + return + else: + leaves = [command] + opening_shapes = { + "ActivityScheduled": "activity", "ActivityStarted": "activity", + "TimerScheduled": "timer", "ChildWorkflowScheduled": "child workflow", + "ChildRunStarted": "child workflow", "ConditionWaitOpened": "condition wait", + } + for offset, leaf in enumerate(leaves): + sequence = boundary.sequence + offset + for event in events: + event_type = _history_event_type(event) + payload = event.get("payload") or {} + if _workflow_sequence(payload) != sequence or event_type not in opening_shapes: + continue + if event_type == "TimerScheduled" and _internal_timeout_timer_kind( + payload, condition_wait_ids_by_sequence, + ) is not None: + continue + _assert_step_matches(leaf, _RecordedStep( + workflow_sequence=sequence, shape=opening_shapes[event_type], + event_types=[event_type], details=_recorded_step_details(payload), + )) + def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: _apply_due_receivers() _consume_terminal_condition_reopens() @@ -5044,6 +5199,36 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: _assert_no_unconsumed_history("complete workflow") return _state(commands + [CompleteWorkflow(result=value)]) + def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _ReplayState | None: + nonlocal cancellation_consumed, authored_sequence, wait_yield_count, advanced_cmd + _assert_cancellation_call_matches(command, boundary) + if boundary.call_kind == "condition": + for event in events: + payload = event.get("payload") or {} + if ( + _history_event_type(event) == "ConditionWaitOpened" + and _workflow_sequence(payload) == boundary.sequence + ): + _apply_condition_wait_receivers(payload.get("condition_wait_id")) + if ( + wait_yield_count < len(recorded_wait_steps) + and recorded_wait_steps[wait_yield_count].workflow_sequence == boundary.sequence + ): + wait_yield_count += 1 + cancellation_consumed = True + ctx._cancel_requested = True + ctx._cancellation_request_id = boundary.request_id + _apply_due_receivers() + if isinstance(command, DurableOperationHandle): + authored_sequence += 1 + try: + advanced_cmd = gen.throw(WorkflowCancelled( + "workflow cancellation was requested", request_id=boundary.request_id, + )) + return None + except StopIteration as stop: + return _terminal_state(stop.value, include_pending=True) + try: while True: # Cursor-0 receivers are start-boundary events. Enter run() once @@ -5062,6 +5247,36 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: return _terminal_state(stop.value, include_pending=True) first = False _apply_due_receivers() + boundary = _cancellation_boundary(cmd) + if cancellation.delivery is not None and not cancellation_consumed: + cancellation_marker = cancellation.delivery + if cancellation_marker.sequence < current_call_sequence: + raise NonDeterministicReplayError( + cancellation_marker.sequence, "committed cancellation boundary", + ["CooperativeCancellationDelivered"], + detail="workflow passed the committed authored call", + ) + if cancellation_marker.sequence == current_call_sequence: + if boundary != cancellation_marker or ctx._cancellation_shield_depth > 0: + raise NonDeterministicReplayError( + cancellation_marker.sequence, "matching unshielded authored call", + ["CooperativeCancellationDelivered"], + detail="committed cancellation call kind or range changed", + ) + terminal = _consume_cancellation(cmd, boundary) + if terminal is not None: + return terminal + continue + elif boundary is not None and not cancellation_consumed and ctx._cancellation_shield_depth == 0: + eligible_sequence = boundary.operation_sequence or boundary.sequence + eligible_span = ( + boundary.operation_sequence_span + if boundary.operation_sequence is not None else boundary.sequence_span + ) + if cancellation.eligible(eligible_sequence, eligible_span): + _assert_cancellation_call_matches(cmd, boundary) + cancellation_intent = boundary + return _state(pending) if isinstance(cmd, SelectGroup): marker = ( selection_markers[selection_marker_cursor] @@ -5374,6 +5589,34 @@ def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: resolution: str | None = None opened: dict[str, Any] | None = None if wait_yield_count < len(wait_opened): + if cancellation.request is not None and wait_yield_count < len(recorded_wait_steps): + physical_sequence = recorded_wait_steps[wait_yield_count].workflow_sequence + if physical_sequence > current_call_sequence: + current_call_sequence = physical_sequence + authored_sequence = max(authored_sequence, physical_sequence + 1) + physical_boundary = CancellationDelivery( + cancellation.request.request_id, physical_sequence, "condition", + ) + if ( + cancellation.delivery is not None and not cancellation_consumed + and cancellation.delivery.sequence == physical_sequence + ): + if cancellation.delivery != physical_boundary or ctx._cancellation_shield_depth > 0: + raise NonDeterministicReplayError( + physical_sequence, "matching unshielded reopened condition", + ["CooperativeCancellationDelivered"], + ) + terminal = _consume_cancellation(cmd, physical_boundary) + if terminal is not None: + return terminal + break + if ( + cancellation.delivery is None and not cancellation_consumed + and ctx._cancellation_shield_depth == 0 and cancellation.eligible(physical_sequence) + ): + _assert_cancellation_call_matches(cmd, physical_boundary) + cancellation_intent = physical_boundary + return _state(pending) if wait_yield_count < len(recorded_wait_steps): _assert_step_matches(cmd, recorded_wait_steps[wait_yield_count]) opened = wait_opened[wait_yield_count] diff --git a/tests/fixtures/replay_regressions/cooperative-reopened-condition-cleanup.json b/tests/fixtures/replay_regressions/cooperative-reopened-condition-cleanup.json new file mode 100644 index 0000000..bb3ac12 --- /dev/null +++ b/tests/fixtures/replay_regressions/cooperative-reopened-condition-cleanup.json @@ -0,0 +1,33 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "cooperative-reopened-condition-cleanup", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "tests.replay.cooperative-reopened-condition-cleanup", + "input": [], + "payload_codec": "avro" + }, + "history": [ + {"event_type": "ConditionWaitOpened", "payload": {"sequence": 1, "condition_key": "forward-wait", "condition_wait_id": "wait-1"}}, + {"event_type": "ConditionWaitSatisfied", "payload": {"sequence": 1, "condition_wait_id": "wait-1"}}, + {"event_type": "ConditionWaitOpened", "payload": {"sequence": 2, "condition_key": "forward-wait", "condition_wait_id": "wait-2"}}, + { + "event_type": "CooperativeCancellationRequested", + "workflow_command_id": "request-1", + "recorded_at": "2026-09-30T12:00:00Z", + "payload": {"workflow_command_id": "request-1", "workflow_run_id": "run-1", "cleanup_deadline_at": "2026-09-30T12:10:00Z"} + }, + { + "event_type": "CooperativeCancellationDelivered", + "workflow_command_id": "request-1", + "payload": {"workflow_command_id": "request-1", "workflow_run_id": "run-1", "sequence": 2, "call_kind": "condition"} + }, + {"event_type": "TimerScheduled", "payload": {"sequence": 3, "timer_kind": "durable_timer", "delay_seconds": 1}}, + {"event_type": "TimerFired", "payload": {"sequence": 3, "timer_kind": "durable_timer"}} + ], + "expected": { + "command_sequence": [{"type": "complete_workflow", "command_type": "CompleteWorkflow", "result": "request-1"}] + } +} diff --git a/tests/test_cooperative_cancellation.py b/tests/test_cooperative_cancellation.py index 6d28bc0..4516421 100644 --- a/tests/test_cooperative_cancellation.py +++ b/tests/test_cooperative_cancellation.py @@ -5,8 +5,11 @@ import pytest +from durable_workflow import serializer from durable_workflow._cooperative_cancellation import CancellationDelivery, read_cancellation_history -from durable_workflow.errors import NonDeterministicReplayError +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.workflow import CompleteWorkflow, ScheduleActivity, StartTimer, WorkflowContext, replay +from tests.test_durable_selection import _activity_completed, _activity_scheduled, _winner_marker def observation() -> dict[str, Any]: @@ -162,3 +165,261 @@ def test_delivery_run_and_request_identity_are_checked() -> None: event["payload"][field] = value with pytest.raises(NonDeterministicReplayError): read_cancellation_history([request(), event], run_id="run-1") + + +def completed_activity(sequence: int, name: str, result: Any) -> dict[str, Any]: + return { + "event_type": "ActivityCompleted", + "payload": { + "sequence": sequence, + "activity_type": name, + "result": serializer.envelope(result), + }, + } + + +class PriorResultCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + assert not ctx.is_cancellation_requested + ctx.throw_if_cancellation_requested() + previous = yield ctx.schedule_activity("previous", []) + try: + yield ctx.start_timer(30) + except WorkflowCancelled as exc: + assert ctx.is_cancellation_requested + with ctx.cancellation_shield(): + ctx.throw_if_cancellation_requested() + cleanup = yield ctx.schedule_activity("cleanup", [previous]) + return {"previous": previous, "cleanup": cleanup, "request_id": exc.request_id} + return "not cancelled" + + +def test_observation_and_request_do_not_throw_before_the_authored_call() -> None: + prior = completed_activity(1, "previous", "already committed") + for history in [[prior], [prior, request()]]: + outcome = replay(PriorResultCleanup, history, [], run_id="run-1", cancellation_request=observation()) + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 2, "timer") + + +def test_cold_replay_reaches_same_delivery_and_durable_cleanup_boundary() -> None: + history = [ + completed_activity(1, "previous", "already committed"), + {"event_type": "TimerScheduled", "payload": {"sequence": 2, "timer_kind": "durable_timer"}}, + request(), + {"event_type": "TimerCancelled", "payload": {"sequence": 2, "timer_kind": "durable_timer"}}, + marker(2, "timer"), + ] + for _restart in range(2): + outcome = replay(PriorResultCleanup, history, [], run_id="run-1") + assert outcome.cancellation_delivery is None + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], ScheduleActivity) + assert outcome.commands[0].activity_type == "cleanup" + assert outcome.commands[0].arguments == ["already committed"] + + history.extend( + [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 3, "activity_type": "cleanup"}}, + completed_activity(3, "cleanup", "cleanup committed"), + ] + ) + for _restart in range(2): + outcome = replay(PriorResultCleanup, history, [], run_id="run-1") + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], CompleteWorkflow) + assert outcome.commands[0].result == { + "previous": "already committed", + "cleanup": "cleanup committed", + "request_id": "request-1", + } + + +class ScalarCleanup: + def run(self, ctx: WorkflowContext, kind: str): # type: ignore[no-untyped-def] + call = { + "activity": lambda: ctx.schedule_activity("forward", []), + "local_activity": lambda: ctx.local_activity("forward", []), + "timer": lambda: ctx.start_timer(30), + "child": lambda: ctx.start_child_workflow("forward-child", []), + "condition": lambda: ctx.wait_condition(lambda: False, key="forward-wait"), + }[kind]() + try: + yield call + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return exc.request_id + return "not cancelled" + + +@pytest.mark.parametrize("kind", ["activity", "local_activity", "timer", "child", "condition"]) +def test_scalar_marker_interrupts_at_call_and_cleanup_does_not_redeliver(kind: str) -> None: + calls: list[Any] = [] + history = [request(), marker(1, kind)] + outcome = replay(ScalarCleanup, history, [kind], run_id="run-1", local_activity_executor=calls.append) + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], StartTimer) + assert outcome.commands[0].delay_seconds == 1 + assert calls == [] + outcome = replay( + ScalarCleanup, + history + + [ + {"event_type": "TimerFired", "payload": {"sequence": 2, "timer_kind": "durable_timer"}}, + ], + [kind], + run_id="run-1", + ) + assert outcome.commands == [CompleteWorkflow("request-1")] + + +def test_pending_local_call_is_delivered_before_executing_the_callable() -> None: + calls: list[Any] = [] + outcome = replay( + ScalarCleanup, [request()], ["local_activity"], run_id="run-1", local_activity_executor=calls.append + ) + assert calls == [] + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 1, "local_activity") + + +class ParallelCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield [ctx.schedule_activity("first", []), [ctx.start_timer(30), ctx.start_child_workflow("child", [])]] + except WorkflowCancelled: + with ctx.cancellation_shield(): + yield ctx.schedule_activity("cleanup", []) + return "cleaned" + + +def test_parallel_delivery_preserves_flat_span_and_cleanup_cursor_on_restart() -> None: + observed = replay(ParallelCleanup, [request()], [], run_id="run-1") + assert observed.cancellation_delivery == CancellationDelivery("request-1", 1, "parallel", 3) + history = [request(), marker(1, "parallel", sequence_span=3), completed_activity(4, "cleanup", "done")] + assert replay(ParallelCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow("cleaned")] + + +@pytest.mark.parametrize( + "history,kind", + [ + ([request(), marker(1, "timer")], "activity"), + ([{"event_type": "TimerFired", "payload": {"sequence": 1}}, request(), marker(2, "activity")], "timer"), + ([request(), marker(1, "parallel", sequence_span=2)], "timer"), + ], +) +def test_changed_or_unreached_delivery_boundary_is_non_deterministic(history: list[dict[str, Any]], kind: str) -> None: + with pytest.raises(NonDeterministicReplayError): + replay(ScalarCleanup, history, [kind], run_id="run-1") + + +class ShieldBeforeRequest: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + try: + yield ctx.start_timer(2) + except WorkflowCancelled as exc: + return exc.request_id + + +def test_shield_defers_first_delivery_without_resetting_request_identity() -> None: + first = replay(ShieldBeforeRequest, [request()], [], run_id="run-1") + assert first.commands == [StartTimer(1)] + assert first.cancellation_delivery is None + history = [request(), {"event_type": "TimerFired", "payload": {"sequence": 1}}] + second = replay(ShieldBeforeRequest, history, [], run_id="run-1") + assert second.cancellation_delivery == CancellationDelivery("request-1", 2, "timer") + cold = replay(ShieldBeforeRequest, history + [marker(2, "timer")], [], run_id="run-1") + assert cold.commands == [CompleteWorkflow("request-1")] + + +class SelectionCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + selected = yield ctx.select({ + "slow": ctx.schedule_activity("slow-activity", []), + "fast": ctx.schedule_activity("fast-activity", []), + }) + yield selected.handles["slow"].await_result() + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return exc.request_id + + +def selection_history() -> list[dict[str, Any]]: + return [ + _activity_scheduled(0, "slow"), _activity_scheduled(1, "fast"), + _activity_completed(1, "fast", "winner"), _winner_marker(), + ] + + +def test_resolved_selection_replays_before_pending_loser_delivery() -> None: + outcome = replay(SelectionCleanup, selection_history() + [request()], [], run_id="run-1") + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 3, "selection_handle", 1, 1) + + +def test_selection_handle_cold_delivery_preserves_opening_identity_and_cleanup_cursor() -> None: + history = selection_history() + [request(), marker(3, "selection_handle", operation_sequence=1)] + for _restart in range(2): + outcome = replay(SelectionCleanup, history, [], run_id="run-1") + assert outcome.commands == [StartTimer(1)] + history.append({"event_type": "TimerFired", "payload": {"sequence": 4}}) + assert replay(SelectionCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + + +def test_first_selection_call_delivers_as_one_parallel_boundary() -> None: + outcome = replay(SelectionCleanup, [request()], [], run_id="run-1") + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 1, "parallel", 2) + cold = replay(SelectionCleanup, [request(), marker(1, "parallel", sequence_span=2)], [], run_id="run-1") + assert cold.commands == [StartTimer(1)] + + +def test_marker_cannot_replace_a_result_committed_before_request() -> None: + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([completed_activity(1, "previous", "result"), request(), marker(1)]) + + +class ReopenedConditionCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.wait_condition(lambda: False, key="forward-wait") + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.schedule_activity("cleanup", []) + return exc.request_id + + +def reopened_condition_history() -> list[dict[str, Any]]: + return [ + {"event_type": "ConditionWaitOpened", "payload": { + "sequence": 1, "condition_key": "forward-wait", "condition_wait_id": "wait-1", + }}, + {"event_type": "ConditionWaitSatisfied", "payload": { + "sequence": 1, "condition_wait_id": "wait-1", + }}, + {"event_type": "ConditionWaitOpened", "payload": { + "sequence": 2, "condition_key": "forward-wait", "condition_wait_id": "wait-2", + }}, + request(), + ] + + +def test_request_during_a_false_condition_reopen_uses_the_actual_wait_occurrence() -> None: + outcome = replay(ReopenedConditionCleanup, reopened_condition_history(), [], run_id="run-1") + assert outcome.commands == [] + assert outcome.cancellation_delivery == CancellationDelivery("request-1", 2, "condition") + + +def test_cold_delivery_during_a_false_reopen_consumes_the_interrupted_wait_once() -> None: + history = reopened_condition_history() + [marker(2, "condition")] + for _restart in range(2): + outcome = replay(ReopenedConditionCleanup, history, [], run_id="run-1") + assert len(outcome.commands) == 1 + assert isinstance(outcome.commands[0], ScheduleActivity) + assert outcome.commands[0].activity_type == "cleanup" + history.append(completed_activity(3, "cleanup", "done")) + assert replay(ReopenedConditionCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] diff --git a/tests/test_replay_regression_corpus.py b/tests/test_replay_regression_corpus.py index 52d1af0..9875162 100644 --- a/tests/test_replay_regression_corpus.py +++ b/tests/test_replay_regression_corpus.py @@ -10,7 +10,7 @@ from durable_workflow import Replayer, Worker, serializer, workflow from durable_workflow.client import Client, WorkflowStreamAppendItem -from durable_workflow.errors import NonDeterministicReplayError, WorkflowPayloadDecodeError +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled, WorkflowPayloadDecodeError from durable_workflow.workflow import WorkflowContext, commands_to_server_commands, query_state from tests.test_golden_history_replay import ( GoldenSagaCompensationWorkflow, @@ -212,8 +212,21 @@ def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] return (yield ctx.local_activity("golden.local", [])) +@workflow.defn(name="tests.replay.cooperative-reopened-condition-cleanup") +class CooperativeReopenedConditionCleanupWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.wait_condition(lambda: False, key="forward-wait") + except WorkflowCancelled as exc: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return exc.request_id + return "not cancelled" + + WORKFLOWS = [ ColdReplacementSatisfiedConditionWorkflow, + CooperativeReopenedConditionCleanupWorkflow, GoldenSagaCompensationWorkflow, GoldenSignalWaitWorkflow, GoldenSingleActivityWorkflow, From a3c909c403634061cea81ebeb200342f31f14d9e Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Wed, 30 Sep 2026 23:57:11 +0000 Subject: [PATCH 03/41] Add fenced cooperative cancellation request and delivery client APIs --- src/durable_workflow/client.py | 163 ++++++++++- tests/test_cooperative_cancellation_client.py | 277 ++++++++++++++++++ 2 files changed, 437 insertions(+), 3 deletions(-) create mode 100644 tests/test_cooperative_cancellation_client.py diff --git a/src/durable_workflow/client.py b/src/durable_workflow/client.py index a293f2c..a46e481 100644 --- a/src/durable_workflow/client.py +++ b/src/durable_workflow/client.py @@ -37,11 +37,13 @@ import httpx from . import serializer +from ._cooperative_cancellation import CancellationDelivery, CancellationRequest from .errors import ( ExternalPayloadIntegrityMismatch, ExternalPayloadOversized, ExternalPayloadUnavailable, ExternalPayloadUnsupported, + NonDeterministicReplayError, RuntimeCapabilityUnsupported, RuntimeDiscoveryUnavailable, ServerError, @@ -87,6 +89,7 @@ CONTROL_PLANE_REQUEST_CONTRACT_SCHEMA = "durable-workflow.v2.control-plane-request.contract" CONTROL_PLANE_REQUEST_CONTRACT_VERSION = 1 _QUERY_TASKS_DISCOVERY_PATH = "worker_protocol.server_capabilities.query_tasks" +_COOPERATIVE_CANCELLATION_DISCOVERY_PATH = "worker_protocol.server_capabilities.cooperative_cancellation" _UPDATE_WAIT_STAGES_DISCOVERY_PATH = ( "control_plane.request_contract.operations.update.fields.wait_for.canonical_values" ) @@ -183,6 +186,13 @@ def _worker_protocol_supports_message_streams() -> bool: return (int(parts[0]), int(parts[1])) >= _MESSAGE_STREAMS_MINIMUM_WORKER_PROTOCOL +def _supports_cooperative_cancellation_protocol(version: Any) -> bool: + if not isinstance(version, str): + return False + match = re.fullmatch(r"1\.([0-9]{1,4})", version) + return match is not None and int(match.group(1)) >= 20 + + def _normalize_base_url(base_url: str) -> str: parsed = urlsplit(base_url) if parsed.query or parsed.fragment: @@ -1330,6 +1340,15 @@ async def cancel(self, *, reason: str | None = None) -> None: """Close this workflow's current run as cancelled. See :meth:`Client.cancel_workflow`.""" await self._client.cancel_workflow(self.workflow_id, reason=reason) + async def request_cancellation( + self, *, reason: str | None = None, cleanup_timeout_seconds: int | None = None, + ) -> dict[str, Any]: + """Request bounded cleanup for this run. See :meth:`Client.request_workflow_cancellation`.""" + return await self._client.request_workflow_cancellation( + self.workflow_id, run_id=self.run_id, reason=reason, + cleanup_timeout_seconds=cleanup_timeout_seconds, + ) + async def terminate(self, *, reason: str | None = None) -> None: """Forcefully stop this workflow. See :meth:`Client.terminate_workflow`.""" await self._client.terminate_workflow(self.workflow_id, reason=reason) @@ -2522,6 +2541,32 @@ async def _require_query_support(self) -> None: ), ) + async def _require_cooperative_cancellation_support(self) -> None: + operation = "Client.request_workflow_cancellation" + info = await self._runtime_discovery( + operation=operation, + required_path=_COOPERATIVE_CANCELLATION_DISCOVERY_PATH, + ) + protocol = info.get("worker_protocol") + capabilities = protocol.get("server_capabilities") if isinstance(protocol, dict) else None + supported = capabilities.get("cooperative_cancellation") if isinstance(capabilities, dict) else None + if supported is False: + raise RuntimeCapabilityUnsupported( + operation, _COOPERATIVE_CANCELLATION_DISCOVERY_PATH, + "This runtime does not support cooperative cancellation. " + "Existing cancel and terminate close immediately.", + ) + if supported is not True: + raise RuntimeDiscoveryUnavailable( + operation, _COOPERATIVE_CANCELLATION_DISCOVERY_PATH, + "Runtime discovery did not advertise cooperative cancellation support.", + ) + if not isinstance(protocol, dict) or not _supports_cooperative_cancellation_protocol(protocol.get("version")): + raise RuntimeDiscoveryUnavailable( + operation, "worker_protocol.version", + "Cooperative cancellation requires an advertised compatible worker protocol of at least 1.20.", + ) + async def _require_update_wait_stage(self, wait_for: str) -> None: info = await self._runtime_discovery( operation="Client.update_workflow", @@ -4140,14 +4185,61 @@ async def query_workflow( "POST", f"/workflows/{workflow_id}/query/{query_name}", json=body, context=workflow_id ) + async def request_workflow_cancellation( + self, + workflow_id: str, + *, + run_id: str | None = None, + reason: str | None = None, + cleanup_timeout_seconds: int | None = None, + ) -> dict[str, Any]: + """Request cancellation with bounded workflow-authored cleanup. + + Requires explicit runtime capability discovery and protocol 1.20. + Repeated requests return Server's original request ID and cleanup + deadline. No client-generated identity or deadline replaces them. + ``run_id`` fences the request to that current run when supplied. + """ + if not isinstance(workflow_id, str) or not workflow_id.strip(): + raise ValueError("workflow_id must be a non-empty string") + if run_id is not None and (not isinstance(run_id, str) or not run_id.strip()): + raise ValueError("run_id must be a non-empty string") + if cleanup_timeout_seconds is not None and ( + type(cleanup_timeout_seconds) is not int or not 1 <= cleanup_timeout_seconds <= 3600 + ): + raise ValueError("cleanup_timeout_seconds must be an integer from 1 to 3600") + await self._require_cooperative_cancellation_support() + path = f"/workflows/{quote(workflow_id, safe='._:-')}" + if run_id is not None: + path += f"/runs/{quote(run_id, safe='._:-')}" + body: dict[str, Any] = {} + if reason is not None: + body["reason"] = reason + if cleanup_timeout_seconds is not None: + body["cleanup_timeout_seconds"] = cleanup_timeout_seconds + result = await self._request("POST", f"{path}/request-cancellation", json=body, context=workflow_id) + if ( + not isinstance(result, dict) or result.get("accepted") is not True + or type(result.get("duplicate")) is not bool or result.get("workflow_id") != workflow_id + or not isinstance(result.get("run_id"), str) or not result["run_id"].strip() + or (run_id is not None and result["run_id"] != run_id) + or not isinstance(result.get("cancellation_request"), dict) + ): + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_response"}) + try: + CancellationRequest.from_observation(result["cancellation_request"]) + except NonDeterministicReplayError as error: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_response"}) from error + return result + async def cancel_workflow(self, workflow_id: str, *, reason: str | None = None) -> None: """Close the current run as cancelled immediately. Server cancels open tasks and timers; it does not resume workflow code to run saga or ``finally`` cleanup. :meth:`terminate_workflow` also - closes immediately, with a distinct terminal outcome. Embedded - Laravel's cooperative ``requestCancellation()`` is not yet available - through this service-mode API. + closes immediately, with a distinct terminal outcome. + For capable runtimes, :meth:`request_workflow_cancellation` separately + requests bounded workflow-authored cleanup. """ body: dict[str, Any] = {} if reason is not None: @@ -5053,6 +5145,71 @@ async def heartbeat_workflow_task( }, ) + async def deliver_workflow_cancellation( + self, + *, + task_id: str, + lease_owner: str, + workflow_task_attempt: int, + request_id: str, + sequence: int, + call_kind: str, + sequence_span: int = 1, + operation_sequence: int | None = None, + operation_sequence_span: int = 1, + ) -> dict[str, Any]: + """Commit delivery at one authored call using the current task lease. + + Requires worker protocol 1.20 and a capable recorded claim. Retry + with the same request, attempt and operation range after a lost + acknowledgment, then reload canonical history before cleanup replay. + This method does not release or renew the lease. + """ + version = _protocol_version_from_env("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION) + if not _supports_cooperative_cancellation_protocol(version): + raise ValueError("cooperative cancellation delivery requires worker protocol 1.20 or newer") + if not isinstance(task_id, str) or not task_id.strip(): + raise ValueError("task_id must be a non-empty string") + if not isinstance(lease_owner, str) or not lease_owner.strip(): + raise ValueError("lease_owner must be a non-empty string") + if type(workflow_task_attempt) is not int or workflow_task_attempt < 1: + raise ValueError("workflow_task_attempt must be a positive integer") + try: + delivery = CancellationDelivery.from_payload({ + "workflow_command_id": request_id, "sequence": sequence, "call_kind": call_kind, + "sequence_span": sequence_span, "operation_sequence": operation_sequence, + "operation_sequence_span": operation_sequence_span, + }) + except NonDeterministicReplayError as error: + raise ValueError("cancellation delivery must name a valid authored call and operation range") from error + body: dict[str, Any] = { + "lease_owner": lease_owner, "workflow_task_attempt": workflow_task_attempt, + "request_id": delivery.request_id, "sequence": delivery.sequence, + "call_kind": delivery.call_kind, "sequence_span": delivery.sequence_span, + } + if delivery.operation_sequence is not None: + body["operation_sequence"] = delivery.operation_sequence + body["operation_sequence_span"] = delivery.operation_sequence_span + result = await self._request( + "POST", f"/worker/workflow-tasks/{quote(task_id, safe='._:-')}/deliver-cancellation", + worker=True, json=body, + ) + if not isinstance(result, dict) or result.get("delivered") is not True or result.get("task_id") != task_id: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_delivery"}) + try: + recorded = CancellationDelivery.from_payload({ + "workflow_command_id": result.get("request_id"), + "sequence": result.get("sequence"), "call_kind": result.get("call_kind"), + "sequence_span": result.get("sequence_span"), + "operation_sequence": result.get("operation_sequence"), + "operation_sequence_span": result.get("operation_sequence_span"), + }) + except NonDeterministicReplayError as error: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_delivery"}) from error + if recorded != delivery: + raise ServerError(200, {"reason": "invalid_cooperative_cancellation_delivery"}) + return result + async def workflow_task_history( self, *, diff --git a/tests/test_cooperative_cancellation_client.py b/tests/test_cooperative_cancellation_client.py new file mode 100644 index 0000000..8364260 --- /dev/null +++ b/tests/test_cooperative_cancellation_client.py @@ -0,0 +1,277 @@ +from __future__ import annotations + +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +import pytest_asyncio + +from durable_workflow.client import Client, WorkflowHandle +from durable_workflow.errors import RuntimeCapabilityUnsupported, RuntimeDiscoveryUnavailable, ServerError +from durable_workflow.retry_policy import TransportRetryPolicy + + +def request_response(*, duplicate: bool = False, run_id: str = "run-1") -> dict[str, Any]: + return { + "accepted": True, "duplicate": duplicate, "workflow_id": "wf/1", "run_id": run_id, + "cancellation_request": { + "request_id": "original-request", "requested_at": "2026-09-30T12:00:00Z", + "cleanup_deadline_at": "2026-09-30T12:10:00Z", "history_refresh_page_token": "opaque-first-page", + "delivery_sequence": None, "delivered_at": None, + }, + } + + +def delivery_response(**overrides: Any) -> dict[str, Any]: + return { + "delivered": True, "task_id": "task/1", "workflow_run_id": "run-1", + "request_id": "original-request", "sequence": 3, "call_kind": "activity", "sequence_span": 1, + "operation_sequence": None, "operation_sequence_span": 1, "reason": None, **overrides, + } + + +def response(body: Any, status: int = 200) -> httpx.Response: + return httpx.Response(status, json=body, request=httpx.Request("POST", "http://test")) + + +@pytest_asyncio.fixture +async def client() -> AsyncIterator[Client]: + async with Client( + "http://localhost:8080", control_token="control-token", worker_token="worker-token", namespace="ns1", + retry_policy=TransportRetryPolicy(initial_backoff_seconds=0, jitter=False), + ) as value: + value._cluster_info = {"worker_protocol": { + "version": "1.20", "server_capabilities": {"cooperative_cancellation": True}, + }} + yield value + + +@pytest.mark.asyncio +@pytest.mark.parametrize("run_id", [None, "run/1"]) +async def test_request_uses_control_credential_and_exact_run_route(client: Client, run_id: str | None) -> None: + expected = request_response(run_id=run_id or "run-1") + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=response(expected, 202)) as send: + result = await client.request_workflow_cancellation( + "wf/1", run_id=run_id, reason="stop", cleanup_timeout_seconds=30, + ) + suffix = "/runs/run%2F1" if run_id is not None else "" + assert send.call_args.args[:2] == ("POST", f"/api/workflows/wf%2F1{suffix}/request-cancellation") + assert send.call_args.kwargs["headers"]["Authorization"] == "Bearer control-token" + assert send.call_args.kwargs["headers"]["X-Namespace"] == "ns1" + assert send.call_args.kwargs["headers"]["X-Durable-Workflow-Control-Plane-Version"] == "2" + assert send.call_args.kwargs["json"] == {"reason": "stop", "cleanup_timeout_seconds": 30} + assert result == expected + + +@pytest.mark.asyncio +async def test_duplicate_returns_original_server_identity_and_deadline(client: Client) -> None: + first = request_response() + duplicate = request_response(duplicate=True) + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[response(first, 202), response(duplicate)]) as send: + accepted = await client.request_workflow_cancellation("wf/1") + repeated = await client.request_workflow_cancellation("wf/1", cleanup_timeout_seconds=3600) + assert accepted["cancellation_request"] == repeated["cancellation_request"] + assert repeated["duplicate"] is True + assert send.await_args_list[0].kwargs["json"] == {} + assert send.await_args_list[1].kwargs["json"] == {"cleanup_timeout_seconds": 3600} + + +@pytest.mark.asyncio +async def test_run_bound_handle_preserves_selection(client: Client) -> None: + with patch.object(client, "request_workflow_cancellation", new_callable=AsyncMock, + return_value=request_response()) as send: + result = await WorkflowHandle(client, "wf/1", "run-1").request_cancellation( + reason="stop", cleanup_timeout_seconds=20, + ) + send.assert_awaited_once_with("wf/1", run_id="run-1", reason="stop", cleanup_timeout_seconds=20) + assert result["cancellation_request"]["request_id"] == "original-request" + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("info", "error"), [ + ({"worker_protocol": {"server_capabilities": {"cooperative_cancellation": False}}}, RuntimeCapabilityUnsupported), + ({}, RuntimeDiscoveryUnavailable), + ({"worker_protocol": {"version": "1.20", "server_capabilities": {"cooperative_cancellation": 1}}}, + RuntimeDiscoveryUnavailable), +]) +async def test_unsupported_or_missing_discovery_sends_no_request(client: Client, info: dict, error: type) -> None: + client._cluster_info = info + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(error): + await client.request_workflow_cancellation("wf/1") + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("version", [None, True, "1.19", "2.20", "one", "1.20", "1." + "2" * 5000]) +async def test_capability_without_compatible_protocol_is_not_admitted(client: Client, version: Any) -> None: + client._cluster_info["worker_protocol"]["version"] = version + with ( + patch.object(client._http, "request", new_callable=AsyncMock) as send, + pytest.raises(RuntimeDiscoveryUnavailable), + ): + await client.request_workflow_cancellation("wf/1") + send.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_first_request_discovers_before_mutation(client: Client) -> None: + info = client._cluster_info + client._cluster_info = None + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[response(info), response(request_response(), 202)]) as send: + await client.request_workflow_cancellation("wf/1") + assert send.await_args_list[0].args[:2] == ("GET", "/api/cluster/info") + assert send.await_args_list[0].kwargs["headers"]["Authorization"] == "Bearer worker-token" + assert send.await_args_list[1].kwargs["headers"]["Authorization"] == "Bearer control-token" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("timeout", [True, 0, -1, 3601, 1.5]) +async def test_invalid_cleanup_timeout_is_rejected_before_io(client: Client, timeout: Any) -> None: + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.request_workflow_cancellation("wf/1", cleanup_timeout_seconds=timeout) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"workflow_id": "other"}, {"run_id": "other"}, {"accepted": 1}, {"duplicate": None}, + {"cancellation_request": None}, + {"cancellation_request": {"request_id": "original-request"}}, + {"cancellation_request": { + **request_response()["cancellation_request"], "cleanup_deadline_at": "2026-09-30T11:00:00Z", + }}, +]) +async def test_malformed_or_mismatched_request_ack_is_rejected(client: Client, change: dict) -> None: + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({**request_response(), **change}, 202)), + pytest.raises(ServerError) as error, + ): + await client.request_workflow_cancellation("wf/1", run_id="run-1") + assert error.value.reason() == "invalid_cooperative_cancellation_response" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selection", [False, True]) +async def test_delivery_sends_owner_attempt_and_authored_boundary( + client: Client, monkeypatch: pytest.MonkeyPatch, selection: bool, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options: dict[str, Any] = {"call_kind": "activity"} + if selection: + options.update(call_kind="selection_handle", operation_sequence=1, operation_sequence_span=2) + with patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(delivery_response(**options))) as send: + result = await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, **options, + ) + assert result["delivered"] is True + assert send.call_args.args[:2] == ("POST", "/api/worker/workflow-tasks/task%2F1/deliver-cancellation") + assert send.call_args.kwargs["headers"]["Authorization"] == "Bearer worker-token" + assert send.call_args.kwargs["headers"]["X-Durable-Workflow-Protocol-Version"] == "1.20" + assert send.call_args.kwargs["json"] == { + "lease_owner": "worker-1", "workflow_task_attempt": 2, "request_id": "original-request", + "sequence": 3, "call_kind": "activity", "sequence_span": 1, **options, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("version", ["1.19", "2.20", "malformed"]) +async def test_delivery_requires_explicit_compatible_protocol( + client: Client, monkeypatch: pytest.MonkeyPatch, version: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", version) + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"workflow_task_attempt": True}, {"sequence": True}, {"call_kind": "made_up"}, + {"sequence_span": 2}, {"call_kind": "selection_handle", "operation_sequence": 3}, + {"operation_sequence": 1}, {"lease_owner": ""}, +]) +async def test_invalid_delivery_is_rejected_before_io( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options = {"task_id": "task/1", "lease_owner": "worker-1", "workflow_task_attempt": 2, + "request_id": "original-request", "sequence": 3, "call_kind": "activity", **change} + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.deliver_workflow_cancellation(**options) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"task_id": "other"}, {"request_id": "other"}, {"sequence": 4}, {"sequence": True}, + {"sequence_span": True}, {"operation_sequence_span": True}, {"call_kind": "timer"}, {"delivered": 1}, +]) +async def test_delivery_ack_must_match_exact_canonical_boundary( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(delivery_response(**change))), pytest.raises(ServerError) as error: + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + assert error.value.reason() == "invalid_cooperative_cancellation_delivery" + + +@pytest.mark.asyncio +async def test_delivery_retries_identical_request_after_acknowledgment_loss( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[httpx.ReadTimeout("acknowledgment lost"), response(delivery_response())]) as send: + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + assert send.await_count == 2 + assert send.await_args_list[0] == send.await_args_list[1] + + +@pytest.mark.asyncio +async def test_active_claim_refusal_does_not_fall_back_to_immediate_cancel(client: Client) -> None: + reason = "active_claim_cancellation_not_supported" + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({"reason": reason}, 409)) as send, + pytest.raises(ServerError) as error, + ): + await client.request_workflow_cancellation("wf/1") + assert error.value.reason() == reason + assert send.await_count == 1 + assert send.call_args.args[1].endswith("/request-cancellation") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reason", ["lease_owner_mismatch", "workflow_task_attempt_mismatch", "lease_expired"]) +async def test_delivery_refusal_preserves_the_server_lease_reason( + client: Client, monkeypatch: pytest.MonkeyPatch, reason: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({"delivered": False, "reason": reason}, 409)) as send, + pytest.raises(ServerError) as error, + ): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="activity", + ) + assert error.value.reason() == reason + assert send.await_count == 1 From febb5906541f4df82cd3f4f3e2c11f41e20acff6 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 00:39:39 +0000 Subject: [PATCH 04/41] Connect cooperative cancellation delivery to Worker replay and local attempt fences --- src/durable_workflow/worker.py | 223 ++++++++-- tests/test_cooperative_cancellation_worker.py | 381 ++++++++++++++++++ 2 files changed, 569 insertions(+), 35 deletions(-) create mode 100644 tests/test_cooperative_cancellation_worker.py diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index e639a3d..28ef1da 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -34,6 +34,7 @@ from typing import Annotated, Any, Concatenate, Literal, ParamSpec, TypeVar, Union, get_args, get_origin, get_type_hints from . import serializer +from ._cooperative_cancellation import CancellationRequest, read_cancellation_history from .activity import ActivityContext, ActivityInfo, _set_context from .auth_composition import ( AUTH_COMPOSITION_CONTRACT_SCHEMA, @@ -49,6 +50,8 @@ PROTOCOL_VERSION, Client, WorkflowExecution, + _protocol_version_from_env, + _supports_cooperative_cancellation_protocol, ) from .errors import ( ActivityCancelled, @@ -87,6 +90,7 @@ NexusServiceCall, RecordLocalActivity, RecordSideEffect, + ReplayOutcome, UpsertMemo, apply_update, commands_to_server_commands, @@ -159,6 +163,10 @@ def __init__(self, kind: str) -> None: self.kind = kind +class _CooperativeCancellationObserved(LocalActivityExecutionAborted): + """Return transport observation to the worker, never to authored cleanup.""" + + class _InvalidLocalActivityReport(NonRetryableError): pass @@ -1021,6 +1029,7 @@ def __init__( } self.activities = {_activity_name(a): a for a in activities} self.capabilities = tuple(dict.fromkeys(capability.strip() for capability in capabilities)) + self._cooperative_cancellation_supported = False if any(not capability for capability in self.capabilities): raise ValueError("worker capabilities must be non-empty strings") self.worker_id = worker_id or f"py-worker-{uuid.uuid4().hex[:8]}" @@ -1177,6 +1186,20 @@ async def _register(self) -> None: raise RuntimeError(f"Server compatibility error: unable to read /api/cluster/info: {e}") from e _validate_server_compatibility(info) + protocol = info.get("worker_protocol") + server_capabilities = protocol.get("server_capabilities") if isinstance(protocol, Mapping) else None + self._cooperative_cancellation_supported = ( + "cooperative_cancellation" in self.capabilities + and isinstance(protocol, Mapping) + and isinstance(server_capabilities, Mapping) + and server_capabilities.get("cooperative_cancellation") is True + and _supports_cooperative_cancellation_protocol(protocol.get("version")) + and _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )) + ) + if "cooperative_cancellation" in self.capabilities and not self._cooperative_cancellation_supported: + raise RuntimeError("cooperative cancellation requires explicit compatible runtime and worker protocol 1.20") self._query_tasks_supported = _server_supports_query_tasks(info) self._workflow_memo_updates_supported = _server_supports_workflow_memo_updates(info) has_update_validators = any( @@ -1389,6 +1412,120 @@ async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None: response.get("renewed") is not True, )): raise LocalActivityExecutionAborted("workflow task lease renewal was not fenced and acknowledged") + observation = response.get("cancellation_request") + if observation is not None: + self._observe_workflow_cancellation(task, observation) + if task.get("_delivered_cancellation_request_id") != task["cancellation_request"]["request_id"]: + raise _CooperativeCancellationObserved("cooperative request observed on the actual task heartbeat") + + def _observe_workflow_cancellation(self, task: dict[str, Any], observation: Any) -> CancellationRequest: + if not self._cooperative_cancellation_supported or not isinstance(observation, Mapping): + raise LocalActivityExecutionAborted("workflow cancellation observation is not negotiated") + try: + current = CancellationRequest.from_observation(observation) + previous = task.get("cancellation_request") + if previous is not None and not isinstance(previous, Mapping): + raise ValueError("previous observation is not an object") + original = CancellationRequest.from_observation(previous) if previous is not None else None + except Exception as error: + raise LocalActivityExecutionAborted("workflow cancellation observation is malformed") from error + if original is not None and (original.request_id, original.requested_at, original.cleanup_deadline_at) != ( + current.request_id, current.requested_at, current.cleanup_deadline_at, + ): + raise LocalActivityExecutionAborted("workflow cancellation observation changed its original identity") + task["cancellation_request"] = dict(observation) + return current + + async def _load_workflow_claim_history( + self, task: dict[str, Any], *, first_page_token: str | None = None, + ) -> list[dict[str, Any]]: + history = [] if first_page_token is not None else list(task.get("history_events", [])) + token = first_page_token if first_page_token is not None else task.get("next_history_page_token") + seen: set[str] = set() + while token is not None: + if not isinstance(token, str) or not token or token in seen: + raise LocalActivityExecutionAborted("workflow history paging did not advance its opaque token") + seen.add(token) + page = await self.client.workflow_task_history( + task_id=task["task_id"], next_history_page_token=token, + lease_owner=self.worker_id, workflow_task_attempt=task.get("workflow_task_attempt", 1), + ) + if not isinstance(page, Mapping) or not isinstance(page.get("history_events"), list): + raise LocalActivityExecutionAborted("workflow history page was not acknowledged") + if any(not isinstance(event, dict) for event in page["history_events"]): + raise LocalActivityExecutionAborted("workflow history page contains a malformed event") + history.extend(page["history_events"]) + token = page.get("next_history_page_token") + return history + + async def _refresh_cancellation_history( + self, task: dict[str, Any], observed: CancellationRequest, + ) -> list[dict[str, Any]]: + try: + return await self._load_workflow_claim_history(task, first_page_token=observed.history_refresh_page_token) + except Exception as error: + raise LocalActivityExecutionAborted( + "canonical cancellation history could not be loaded on this claim", + ) from error + + async def _replay_workflow_claim( + self, cls: type, task: dict[str, Any], history: list[dict[str, Any]], start_input: list[Any], + *, payload_codec: str | None, execute_local: Callable[[RecordLocalActivity], Any], + ) -> tuple[ReplayOutcome, list[dict[str, Any]]]: + observation = task.get("cancellation_request") + if observation is not None: + observed = self._observe_workflow_cancellation(task, observation) + history = await self._refresh_cancellation_history(task, observed) + for _ in range(3): + state = read_cancellation_history( + history, run_id=task.get("run_id", ""), observation=task.get("cancellation_request"), + ) + if state.request is not None and not self._cooperative_cancellation_supported: + raise LocalActivityExecutionAborted("canonical cancellation requires a negotiated capable worker") + if state.delivery is not None: + task["_delivered_cancellation_request_id"] = state.delivery.request_id + try: + outcome = await asyncio.to_thread( + replay, cls, history, start_input, workflow_id=task.get("workflow_id"), + run_id=task.get("run_id", ""), + workflow_command_id=( + _string_or_none(task.get("workflow_command_id")) or _string_or_none(task.get("task_id")) + ), + payload_codec=payload_codec, external_storage=self.external_storage, + external_storage_cache=self.external_storage_cache, + cancel_requested=bool(task.get("cancel_requested", False)) and state.request is None, + cancellation_request=task.get("cancellation_request"), local_activity_executor=execute_local, + ) + except _CooperativeCancellationObserved: + observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) + history = await self._refresh_cancellation_history(task, observed) + continue + intent = outcome.cancellation_delivery + if intent is None or outcome.commands: + # Earlier authored commands must commit on this claim first. + # Completion releases the claim; a successor replays their durable results. + return outcome, history + observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) + delivery_error: Exception | None = None + try: + await self.client.deliver_workflow_cancellation( + task_id=task["task_id"], lease_owner=self.worker_id, + workflow_task_attempt=task.get("workflow_task_attempt", 1), + request_id=intent.request_id, sequence=intent.sequence, call_kind=intent.call_kind, + sequence_span=intent.sequence_span, operation_sequence=intent.operation_sequence, + operation_sequence_span=intent.operation_sequence_span, + ) + except Exception as error: + delivery_error = error + history = await self._refresh_cancellation_history(task, observed) + committed = read_cancellation_history( + history, run_id=task.get("run_id", ""), observation=task.get("cancellation_request"), + ) + if committed.delivery != intent: + raise LocalActivityExecutionAborted( + "delivery was not proved by matching canonical history", + ) from delivery_error + raise LocalActivityExecutionAborted("workflow cancellation replay did not converge on its canonical delivery") def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str, Any]: storage = self.external_storage @@ -1401,6 +1538,42 @@ def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str reference = store_external_payload(storage, data, codec=envelope["codec"]) return {"codec": envelope["codec"], "external_storage": reference.to_dict()} + async def _execute_cooperative_local_callable( + self, task: dict[str, Any], command: RecordLocalActivity, handler: Callable[..., Any], + attempt_state: dict[str, Any], + ) -> Any: + async def observe_lease() -> None: + while True: + await asyncio.sleep(min(5.0, self._heartbeat_interval)) + await self._renew_local_workflow_lease(task) + + invocation = asyncio.create_task(self._execute_activity_callable( + task, command.activity_type, tuple(command.arguments), handler, + )) + observation = asyncio.create_task(observe_lease()) + try: + done, _ = await asyncio.wait([invocation, observation], return_when=asyncio.FIRST_COMPLETED) + if observation in done: + await observation # Propagate transport observation or lost lease, never a workflow cancellation. + raise LocalActivityExecutionAborted("local lease observer stopped without an acknowledgment") + return await invocation + finally: + observation.cancel() + with contextlib.suppress(asyncio.CancelledError, Exception): + await observation + if not invocation.done(): + attempt_state["lease_aborted"] = True + invocation.cancel() + + def discard_late_result(future: asyncio.Task[Any]) -> None: + if not future.cancelled(): + future.exception() + + # Python cannot forcibly stop a synchronous thread or a callable + # that suppresses cancellation. Its attempt is fenced and its + # eventual result cannot become a durable command. + invocation.add_done_callback(discard_late_result) + async def _execute_local_activity( self, task: dict[str, Any], @@ -1478,7 +1651,10 @@ def check_boundary( now = time.monotonic() if state["lease_aborted"]: raise LocalActivityExecutionAborted("local activity lost its workflow task lease") - if self._stop.is_set() or task.get("cancel_requested") is True: + if self._stop.is_set() or ( + task.get("cancel_requested") is True and task.get("cancellation_request") is None + and task.get("_delivered_cancellation_request_id") is None + ): raise ActivityCancelled("local activity cancelled") if ( command.heartbeat_timeout is not None @@ -1541,11 +1717,16 @@ async def heartbeat( ) _set_context(ActivityContext(info=info, client=self.client, heartbeat_callback=heartbeat)) try: - result = await self._execute_activity_callable( - task, command.activity_type, tuple(command.arguments), handler, - ) + if self._cooperative_cancellation_supported: + result = await self._execute_cooperative_local_callable(task, command, handler, attempt_state) + else: + result = await self._execute_activity_callable( + task, command.activity_type, tuple(command.arguments), handler, + ) finally: _set_context(None) + if self._cooperative_cancellation_supported: + await self._renew_local_workflow_lease(task) check_boundary() attempts.append({ "attempt_id": attempt_id, @@ -1636,21 +1817,7 @@ async def _run_workflow_task_core(self, task: dict[str, Any]) -> list[dict[str, task_id: str = task["task_id"] attempt: int = task.get("workflow_task_attempt", 1) wf_type: str = task.get("workflow_type", "") - history = task.get("history_events", []) - - # The worker requests bounded history pages when polling. Do not replay - # an incomplete history if fetching a later page fails. - next_page_token = task.get("next_history_page_token") - while next_page_token: - page_data = await self.client.workflow_task_history( - task_id=task_id, - next_history_page_token=next_page_token, - lease_owner=self.worker_id, - workflow_task_attempt=attempt, - ) - if page_data and page_data.get("history_events"): - history.extend(page_data["history_events"]) - next_page_token = page_data.get("next_history_page_token") if page_data else None + history = await self._load_workflow_claim_history(task) start_input: list[Any] = [] codec = task.get("payload_codec") @@ -1773,22 +1940,8 @@ def execute_local(command: RecordLocalActivity) -> Any: return future.result() try: - outcome = await asyncio.to_thread( - replay, - cls, - history, - start_input, - workflow_id=task.get("workflow_id"), - run_id=run_id, - workflow_command_id=( - _string_or_none(task.get("workflow_command_id")) - or _string_or_none(task.get("task_id")) - ), - payload_codec=codec, - external_storage=self.external_storage, - external_storage_cache=self.external_storage_cache, - cancel_requested=bool(task.get("cancel_requested", False)), - local_activity_executor=execute_local, + outcome, history = await self._replay_workflow_claim( + cls, task, history, start_input, payload_codec=codec, execute_local=execute_local, ) except LocalActivityExecutionAborted as e: log.warning("abandoning workflow task %s before local activity commit: %s", task_id, e) diff --git a/tests/test_cooperative_cancellation_worker.py b/tests/test_cooperative_cancellation_worker.py new file mode 100644 index 0000000..ae0d535 --- /dev/null +++ b/tests/test_cooperative_cancellation_worker.py @@ -0,0 +1,381 @@ +from __future__ import annotations + +import asyncio +from copy import deepcopy +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import activity, serializer, workflow +from durable_workflow.client import Client +from durable_workflow.errors import ServerError, WorkflowCancelled +from durable_workflow.worker import Worker +from tests.test_cooperative_cancellation import marker, observation, request +from tests.test_worker import compatible_cluster_info + + +@workflow.defn(name="worker-cancellation") +class CancellationWorkflow: + def run(self, ctx: workflow.WorkflowContext, kind: str): # type: ignore[no-untyped-def] + ctx.throw_if_cancellation_requested() + try: + if kind == "local_activity": + yield ctx.local_activity("work", []) + elif kind == "activity": + yield ctx.schedule_activity("work", []) + elif kind == "parallel": + yield [ctx.start_timer(20), ctx.schedule_activity("work", [])] + else: + yield ctx.start_timer(30) + except WorkflowCancelled as error: + with ctx.cancellation_shield(): + yield ctx.start_timer(1) + return error.request_id + return "not cancelled" + + +@workflow.defn(name="worker-cancellation-prior") +class PriorCommandsWorkflow: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + yield ctx.upsert_memo({"stage": "before-boundary"}) + yield ctx.start_timer(30) + + +@workflow.defn(name="worker-cancellation-local-cleanup") +class LocalCleanupWorkflow: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(30) + except WorkflowCancelled as error: + with ctx.cancellation_shield(): + yield ctx.local_activity("cleanup", [error.request_id]) + return error.request_id + + +def claimed_task(kind: str = "timer", *, observed: bool = True) -> dict[str, Any]: + return { + "task_id": "task-1", "workflow_id": "workflow-1", "run_id": "run-1", + "workflow_type": "worker-cancellation", "workflow_task_attempt": 4, + "payload_codec": "avro", "arguments": serializer.envelope([kind], codec="avro"), + "history_events": [], **({"cancellation_request": observation()} if observed else {}), + } + + +def lease_ack(*, observed: bool = False) -> dict[str, Any]: + return { + "task_id": "task-1", "lease_owner": "cooperative-worker", "workflow_task_attempt": 4, + "renewed": True, **({"cancellation_request": observation()} if observed else {}), + } + + +class ClaimServer: + def __init__(self, *, history: list[dict[str, Any]] | None = None) -> None: + self.client = AsyncMock(spec=Client) + self.history = list(history if history is not None else [request()]) + self.trace: list[str] = [] + self.delivery_error: Exception | None = None + self.commit_delivery = True + self.client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": { + "query_tasks": True, "long_poll_timeout": 30, "cooperative_cancellation": True, + "workflow_memo_updates": True, + }, + }) + self.client.register_worker.return_value = {"registered": True} + self.client.heartbeat_workflow_task.return_value = lease_ack() + self.client.workflow_task_history.side_effect = self.page + self.client.deliver_workflow_cancellation.side_effect = self.deliver + self.client.complete_workflow_task.side_effect = self.complete + + async def page(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs == { + "task_id": "task-1", "next_history_page_token": "opaque-first-page", + "lease_owner": "cooperative-worker", "workflow_task_attempt": 4, + } + self.trace.append("history") + return {"history_events": deepcopy(self.history), "next_history_page_token": None} + + async def deliver(self, **kwargs: Any) -> dict[str, Any]: + self.trace.append("delivery") + if self.commit_delivery: + self.history.append(marker( + kwargs["sequence"], kwargs["call_kind"], sequence_span=kwargs["sequence_span"], + operation_sequence=kwargs["operation_sequence"], + operation_sequence_span=kwargs["operation_sequence_span"], + )) + if self.delivery_error: + raise self.delivery_error + return {"delivered": True} + + async def complete(self, **kwargs: Any) -> dict[str, Any]: + self.trace.append("completion") + return {"outcome": "completed"} + + async def worker(self, monkeypatch: pytest.MonkeyPatch, **kwargs: Any) -> Worker: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + worker = Worker( + self.client, task_queue="queue", worker_id="cooperative-worker", + workflows=kwargs.pop("workflows", [CancellationWorkflow]), + capabilities=["cooperative_cancellation"], **kwargs, + ) + await worker._register() + return worker + + +@pytest.mark.parametrize("kind,span", [("timer", 1), ("activity", 1), ("local_activity", 1), ("parallel", 2)]) +async def test_claim_commits_delivery_and_reloads_before_cleanup( + monkeypatch: pytest.MonkeyPatch, kind: str, span: int, +) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch) + commands = await worker._run_workflow_task(claimed_task(kind)) + assert server.trace == ["history", "delivery", "history", "completion"] + assert commands is not None and len(commands) == 1 + assert commands[0]["type"] == "start_timer" + assert commands[0]["delay_seconds"] == 1 + assert server.client.deliver_workflow_cancellation.await_args.kwargs == { + "task_id": "task-1", "lease_owner": "cooperative-worker", "workflow_task_attempt": 4, + "request_id": "request-1", "sequence": 1, "call_kind": kind, "sequence_span": span, + "operation_sequence": None, "operation_sequence_span": 1, + } + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_lost_delivery_ack_is_resolved_only_from_canonical_history(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + server.delivery_error = TimeoutError("ack lost after commit") + worker = await server.worker(monkeypatch) + commands = await worker._run_workflow_task(claimed_task()) + assert commands is not None and commands[0]["type"] == "start_timer" + assert server.trace == ["history", "delivery", "history", "completion"] + assert server.client.deliver_workflow_cancellation.await_count == 1 + + +async def test_legacy_boolean_cannot_replace_canonical_cooperative_delivery(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch) + task = claimed_task() + task["cancel_requested"] = True + commands = await worker._run_workflow_task(task) + assert commands is not None and commands[0]["type"] == "start_timer" + assert server.trace == ["history", "delivery", "history", "completion"] + + +@pytest.mark.parametrize("error", [None, TimeoutError("not accepted"), ServerError(409, {"reason": "lease_expired"})]) +async def test_ack_without_matching_marker_cannot_run_cleanup_or_complete( + monkeypatch: pytest.MonkeyPatch, error: Exception | None, +) -> None: + server = ClaimServer() + server.commit_delivery = False + server.delivery_error = error + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(claimed_task()) is None + assert server.trace == ["history", "delivery", "history"] + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_claim_loss_during_refresh_cannot_complete_or_fail_stale_attempt(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + server.client.workflow_task_history.side_effect = [ + {"history_events": [request()], "next_history_page_token": None}, + ServerError(409, {"reason": "workflow_task_attempt_mismatch"}), + ] + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(claimed_task()) is None + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_prior_commands_commit_before_a_successor_can_deliver(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch, workflows=[PriorCommandsWorkflow]) + worker._workflow_memo_updates_supported = True + task = claimed_task() + task.update(workflow_type="worker-cancellation-prior", arguments=serializer.envelope([], codec="avro")) + commands = await worker._run_workflow_task(task) + assert commands is not None and [command["type"] for command in commands] == ["upsert_memo"] + server.client.deliver_workflow_cancellation.assert_not_awaited() + assert server.trace == ["history", "completion"] + + +async def test_cold_replacement_replays_marker_without_redelivery(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer(history=[request(), marker(1, "timer")]) + worker = await server.worker(monkeypatch) + task = claimed_task(observed=False) + task["history_events"] = deepcopy(server.history) + commands = await worker._run_workflow_task(task) + assert commands is not None and commands[0]["delay_seconds"] == 1 + server.client.deliver_workflow_cancellation.assert_not_awaited() + assert server.trace == ["completion"] + + +async def test_local_result_after_request_observation_is_not_serialized_or_committed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer(history=[]) + + @activity.defn(name="work") + async def work() -> object: + server.trace.append("local-result") + server.history.append(request()) + return object() # Would fail serialization if the late value escaped the fence. + + server.client.heartbeat_workflow_task.side_effect = [lease_ack(), lease_ack(observed=True)] + worker = await server.worker(monkeypatch, activities=[work]) + commands = await worker._run_workflow_task(claimed_task("local_activity", observed=False)) + assert server.trace == ["local-result", "history", "delivery", "history", "completion"] + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + assert server.client.deliver_workflow_cancellation.await_args.kwargs["call_kind"] == "local_activity" + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_local_heartbeat_observation_does_not_become_an_activity_failure(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer(history=[]) + + @activity.defn(name="work") + async def work() -> str: + server.history.append(request()) + await activity.context().heartbeat() + raise AssertionError("local execution should have returned to canonical replay") + + server.client.heartbeat_workflow_task.side_effect = [lease_ack(), lease_ack(observed=True)] + worker = await server.worker(monkeypatch, activities=[work]) + commands = await worker._run_workflow_task(claimed_task("local_activity", observed=False)) + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_active_local_call_without_user_heartbeat_is_fenced_before_delivery( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer(history=[]) + late_result_discarded = asyncio.Event() + + @activity.defn(name="work") + async def work() -> object: + server.history.append(request()) + server.client.heartbeat_workflow_task.return_value = lease_ack(observed=True) + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + late_result_discarded.set() + return object() + + worker = await server.worker(monkeypatch, activities=[work], heartbeat_interval=0.01) + commands = await asyncio.wait_for( + worker._run_workflow_task(claimed_task("local_activity", observed=False)), timeout=2, + ) + await asyncio.wait_for(late_result_discarded.wait(), timeout=1) + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + assert server.client.deliver_workflow_cancellation.await_args.kwargs["call_kind"] == "local_activity" + server.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("field,value", [ + ("renewed", False), ("lease_owner", "replacement-worker"), + ("workflow_task_attempt", 5), ("task_id", "replacement-task"), +]) +async def test_active_local_lease_refusal_drops_work_without_delivery( + monkeypatch: pytest.MonkeyPatch, field: str, value: Any, +) -> None: + server = ClaimServer(history=[]) + abandoned = asyncio.Event() + + @activity.defn(name="work") + async def work() -> None: + ack = lease_ack(observed=True) + ack[field] = value + server.client.heartbeat_workflow_task.return_value = ack + try: + await asyncio.Event().wait() + finally: + abandoned.set() + + worker = await server.worker(monkeypatch, activities=[work], heartbeat_interval=0.01) + task = claimed_task("local_activity", observed=False) + assert await asyncio.wait_for(worker._run_workflow_task(task), timeout=2) is None + await asyncio.wait_for(abandoned.wait(), timeout=1) + assert "cancellation_request" not in task + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_shielded_local_cleanup_can_heartbeat_the_delivered_request(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + cleanup: list[str] = [] + + @activity.defn(name="cleanup") + async def compensate(request_id: str) -> str: + cleanup.append(request_id) + await activity.context().heartbeat() + return request_id + + server.client.heartbeat_workflow_task.return_value = lease_ack(observed=True) + worker = await server.worker(monkeypatch, workflows=[LocalCleanupWorkflow], activities=[compensate]) + task = claimed_task() + task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro")) + commands = await worker._run_workflow_task(task) + assert cleanup == ["request-1"] + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + assert server.client.deliver_workflow_cancellation.await_count == 1 + + +@pytest.mark.parametrize("value", ["", None, {}, "not-a-page"]) +async def test_invalid_or_missing_refresh_token_never_delivers(monkeypatch: pytest.MonkeyPatch, value: Any) -> None: + server = ClaimServer() + worker = await server.worker(monkeypatch) + task = claimed_task() + task["cancellation_request"]["history_refresh_page_token"] = value + assert await worker._run_workflow_task(task) is None + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + + +async def test_canonical_refresh_rejects_repeated_page_token(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + server.client.workflow_task_history.side_effect = [{ + "history_events": [request()], "next_history_page_token": "opaque-first-page", + }] + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(claimed_task()) is None + assert server.client.workflow_task_history.await_count == 1 + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("version", ["1.19", "2.20", "1.bad"]) +async def test_explicit_cooperative_worker_refuses_incompatible_protocol( + monkeypatch: pytest.MonkeyPatch, version: str, +) -> None: + server = ClaimServer() + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", version) + worker = Worker(server.client, task_queue="queue", capabilities=["cooperative_cancellation"]) + with pytest.raises(RuntimeError, match="protocol 1.20"): + await worker._register() + server.client.register_worker.assert_not_awaited() + + +async def test_current_default_worker_does_not_advertise_cooperative_support(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer() + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + worker = Worker(server.client, task_queue="queue") + await worker._register() + assert worker._cooperative_cancellation_supported is False + assert "cooperative_cancellation" not in server.client.register_worker.await_args.kwargs["capabilities"] + + +async def test_incapable_worker_cannot_execute_a_canonical_cooperative_run(monkeypatch: pytest.MonkeyPatch) -> None: + server = ClaimServer(history=[request(), marker(1, "timer")]) + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + worker = Worker(server.client, task_queue="queue", workflows=[CancellationWorkflow]) + await worker._register() + task = claimed_task(observed=False) + task["history_events"] = deepcopy(server.history) + assert await worker._run_workflow_task(task) is None + server.client.deliver_workflow_cancellation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() From 797e3935d7a738eb91fc98a396a772ced82bd31f Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 01:21:19 +0000 Subject: [PATCH 05/41] Drain cooperative local cleanup and fence it at shutdown deadline --- src/durable_workflow/worker.py | 41 +++- tests/test_cooperative_cancellation_worker.py | 187 +++++++++++++++++- 2 files changed, 220 insertions(+), 8 deletions(-) diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 28ef1da..23fd226 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -17,6 +17,7 @@ import asyncio import contextlib +import contextvars import hashlib import inspect import json @@ -28,6 +29,7 @@ import types import uuid from collections.abc import Awaitable, Callable, Coroutine, Iterable, Mapping +from concurrent.futures import ThreadPoolExecutor from datetime import datetime, timezone from functools import wraps from types import FunctionType @@ -1056,6 +1058,8 @@ def __init__( self.max_concurrent_worker_sessions = max_concurrent_worker_sessions self._worker_sessions: dict[str, WorkerSession] = {} self._stop = asyncio.Event() + self._local_activity_shutdown = asyncio.Event() + self._local_activity_executor: ThreadPoolExecutor | None = None self._wf_semaphore = asyncio.Semaphore(max_concurrent_workflow_tasks) self._act_semaphore = asyncio.Semaphore(max_concurrent_activity_tasks) self._shutdown_timeout = shutdown_timeout @@ -1395,6 +1399,8 @@ async def _resolve_workflow_nexus_commands( raise RuntimeError("workflow yielded too many consecutive Nexus service calls") async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None: + if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set(): + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") task_id = str(task["task_id"]) attempt = int(task.get("workflow_task_attempt", 1)) try: @@ -1405,6 +1411,8 @@ async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None: ) except Exception as exc: raise LocalActivityExecutionAborted("workflow task lease renewal failed") from exc + if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set(): + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") if not isinstance(response, Mapping) or any(( response.get("task_id") != task_id, response.get("lease_owner") != self.worker_id, @@ -1548,16 +1556,22 @@ async def observe_lease() -> None: await self._renew_local_workflow_lease(task) invocation = asyncio.create_task(self._execute_activity_callable( - task, command.activity_type, tuple(command.arguments), handler, + task, command.activity_type, tuple(command.arguments), handler, run_sync_in_thread=True, )) observation = asyncio.create_task(observe_lease()) + shutdown = asyncio.create_task(self._local_activity_shutdown.wait()) try: - done, _ = await asyncio.wait([invocation, observation], return_when=asyncio.FIRST_COMPLETED) + done, _ = await asyncio.wait([invocation, observation, shutdown], return_when=asyncio.FIRST_COMPLETED) + if shutdown in done: + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") if observation in done: await observation # Propagate transport observation or lost lease, never a workflow cancellation. raise LocalActivityExecutionAborted("local lease observer stopped without an acknowledgment") return await invocation finally: + shutdown.cancel() + with contextlib.suppress(asyncio.CancelledError): + await shutdown observation.cancel() with contextlib.suppress(asyncio.CancelledError, Exception): await observation @@ -1651,7 +1665,10 @@ def check_boundary( now = time.monotonic() if state["lease_aborted"]: raise LocalActivityExecutionAborted("local activity lost its workflow task lease") - if self._stop.is_set() or ( + if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set(): + state["lease_aborted"] = True + raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim") + if (self._stop.is_set() and not self._cooperative_cancellation_supported) or ( task.get("cancel_requested") is True and task.get("cancellation_request") is None and task.get("_delivered_cancellation_request_id") is None ): @@ -2320,6 +2337,7 @@ async def _execute_activity_callable( activity_type: str, args: tuple[Any, ...], fn: Callable[..., Any], + *, run_sync_in_thread: bool = False, ) -> Any: context = ActivityInterceptorContext( worker_id=self.worker_id, @@ -2330,7 +2348,17 @@ async def _execute_activity_callable( ) async def call_activity(ctx: ActivityInterceptorContext) -> Any: - result = fn(*ctx.args) + if run_sync_in_thread and not inspect.iscoroutinefunction(fn): + if self._local_activity_executor is None: + # Replay threads wait for local results, so sharing their pool can deadlock. + self._local_activity_executor = ThreadPoolExecutor( + max_workers=self.max_concurrent_workflow_tasks, thread_name_prefix="dw-local-activity", + ) + result = await asyncio.get_running_loop().run_in_executor( + self._local_activity_executor, contextvars.copy_context().run, fn, *ctx.args, + ) + else: + result = fn(*ctx.args) if asyncio.iscoroutine(result): return await result return result @@ -3527,6 +3555,8 @@ async def _shutdown(self) -> None: in_flight, timeout=self._remaining_shutdown_time(deadline), ) + if pending: + self._local_activity_shutdown.set() for t in pending: t.cancel() if pending: @@ -3541,6 +3571,9 @@ async def _shutdown(self) -> None: ) await asyncio.gather(*in_flight, return_exceptions=True) + if self._local_activity_executor is not None: + self._local_activity_executor.shutdown(wait=False, cancel_futures=True) + for session in self._worker_sessions.values(): if not session.active: continue diff --git a/tests/test_cooperative_cancellation_worker.py b/tests/test_cooperative_cancellation_worker.py index ae0d535..4c0cb46 100644 --- a/tests/test_cooperative_cancellation_worker.py +++ b/tests/test_cooperative_cancellation_worker.py @@ -1,6 +1,8 @@ from __future__ import annotations import asyncio +import threading +from concurrent.futures import ThreadPoolExecutor from copy import deepcopy from typing import Any from unittest.mock import AsyncMock @@ -11,6 +13,7 @@ from durable_workflow.client import Client from durable_workflow.errors import ServerError, WorkflowCancelled from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted from tests.test_cooperative_cancellation import marker, observation, request from tests.test_worker import compatible_cluster_info @@ -70,8 +73,13 @@ def lease_ack(*, observed: bool = False) -> dict[str, Any]: class ClaimServer: - def __init__(self, *, history: list[dict[str, Any]] | None = None) -> None: + def __init__( + self, *, history: list[dict[str, Any]] | None = None, + lease_owner: str = "cooperative-worker", workflow_task_attempt: int = 4, + ) -> None: self.client = AsyncMock(spec=Client) + self.lease_owner = lease_owner + self.workflow_task_attempt = workflow_task_attempt self.history = list(history if history is not None else [request()]) self.trace: list[str] = [] self.delivery_error: Exception | None = None @@ -83,7 +91,9 @@ def __init__(self, *, history: list[dict[str, Any]] | None = None) -> None: }, }) self.client.register_worker.return_value = {"registered": True} - self.client.heartbeat_workflow_task.return_value = lease_ack() + self.client.heartbeat_workflow_task.return_value = { + **lease_ack(), "lease_owner": lease_owner, "workflow_task_attempt": workflow_task_attempt, + } self.client.workflow_task_history.side_effect = self.page self.client.deliver_workflow_cancellation.side_effect = self.deliver self.client.complete_workflow_task.side_effect = self.complete @@ -91,7 +101,7 @@ def __init__(self, *, history: list[dict[str, Any]] | None = None) -> None: async def page(self, **kwargs: Any) -> dict[str, Any]: assert kwargs == { "task_id": "task-1", "next_history_page_token": "opaque-first-page", - "lease_owner": "cooperative-worker", "workflow_task_attempt": 4, + "lease_owner": self.lease_owner, "workflow_task_attempt": self.workflow_task_attempt, } self.trace.append("history") return {"history_events": deepcopy(self.history), "next_history_page_token": None} @@ -115,7 +125,7 @@ async def complete(self, **kwargs: Any) -> dict[str, Any]: async def worker(self, monkeypatch: pytest.MonkeyPatch, **kwargs: Any) -> Worker: monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") worker = Worker( - self.client, task_queue="queue", worker_id="cooperative-worker", + self.client, task_queue="queue", worker_id=self.lease_owner, workflows=kwargs.pop("workflows", [CancellationWorkflow]), capabilities=["cooperative_cancellation"], **kwargs, ) @@ -324,6 +334,175 @@ async def compensate(request_id: str) -> str: assert server.client.deliver_workflow_cancellation.await_count == 1 +async def test_worker_stop_drains_shielded_local_cleanup_within_its_grace_period( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer() + entered = asyncio.Event() + release = asyncio.Event() + + @activity.defn(name="cleanup") + async def compensate(request_id: str) -> str: + entered.set() + await release.wait() + await activity.context().heartbeat() + return request_id + + worker = await server.worker( + monkeypatch, workflows=[LocalCleanupWorkflow], activities=[compensate], shutdown_timeout=1, + ) + task = claimed_task() + task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro")) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=1) + stopping = asyncio.create_task(worker.stop()) + await asyncio.sleep(0) + assert worker._stop.is_set() + release.set() + commands = await asyncio.wait_for(execution, timeout=1) + await asyncio.wait_for(stopping, timeout=1) + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + server.client.deliver_workflow_cancellation.assert_awaited_once() + server.client.fail_workflow_task.assert_not_awaited() + server.client.deregister_worker_registration.assert_awaited_once_with("cooperative-worker") + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +async def test_shutdown_timeout_abandons_cleanup_and_replacement_replays_original_delivery( + monkeypatch: pytest.MonkeyPatch, handler_kind: str, +) -> None: + server = ClaimServer() + entered = asyncio.Event() + discarded = asyncio.Event() + loop = asyncio.get_running_loop() + release_thread = threading.Event() + fenced_heartbeats: list[bool] = [] + + @activity.defn(name="cleanup") + async def compensate(request_id: str) -> object: + assert request_id == "request-1" + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + with pytest.raises(LocalActivityExecutionAborted): + await activity.context().heartbeat() + fenced_heartbeats.append(True) + discarded.set() + return object() # A late, unencodable result must never become an activity failure. + + @activity.defn(name="cleanup") + def synchronous_compensation(request_id: str) -> object: + assert request_id == "request-1" + context = activity.context() + assert context.info.worker_id == "cooperative-worker" + loop.call_soon_threadsafe(entered.set) + try: + assert release_thread.wait(timeout=2) + heartbeat = asyncio.run_coroutine_threadsafe(context.heartbeat(), loop) + try: + heartbeat.result(timeout=1) + except LocalActivityExecutionAborted: + fenced_heartbeats.append(True) + else: + fenced_heartbeats.append(False) + return object() + finally: + loop.call_soon_threadsafe(discarded.set) + + worker = await server.worker( + monkeypatch, workflows=[LocalCleanupWorkflow], + activities=[compensate if handler_kind == "async" else synchronous_compensation], shutdown_timeout=0.01, + ) + task = claimed_task() + task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro")) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=1) + original_history = deepcopy(server.history) + heartbeats_before_stop = server.client.heartbeat_workflow_task.await_count + await asyncio.wait_for(worker.stop(), timeout=1) + release_thread.set() + await asyncio.wait_for(discarded.wait(), timeout=1) + assert fenced_heartbeats == [True] + assert server.client.heartbeat_workflow_task.await_count == heartbeats_before_stop + assert execution.cancelled() or execution.result() is None + assert server.history == original_history + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + server.client.deregister_worker_registration.assert_awaited_once_with("cooperative-worker") + + replacement = ClaimServer(history=original_history, lease_owner="replacement-worker", workflow_task_attempt=5) + cleanup_ids: list[str] = [] + + @activity.defn(name="cleanup") + async def resumed_cleanup(request_id: str) -> str: + cleanup_ids.append(request_id) + await activity.context().heartbeat() + return request_id + + successor = await replacement.worker(monkeypatch, workflows=[LocalCleanupWorkflow], activities=[resumed_cleanup]) + reclaimed = deepcopy(task) + reclaimed["workflow_task_attempt"] = 5 + reclaimed["history_events"] = deepcopy(original_history) + commands = await asyncio.wait_for(successor._run_workflow_task(reclaimed), timeout=1) + assert cleanup_ids == ["request-1"] + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + assert replacement.history == original_history + replacement.client.deliver_workflow_cancellation.assert_not_awaited() + replacement.client.fail_workflow_task.assert_not_awaited() + assert replacement.client.complete_workflow_task.await_args.kwargs["lease_owner"] == "replacement-worker" + assert replacement.client.complete_workflow_task.await_args.kwargs["workflow_task_attempt"] == 5 + + +async def test_active_synchronous_local_call_observes_request_without_blocking_the_worker( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = ClaimServer(history=[]) + entered = asyncio.Event() + discarded = asyncio.Event() + release = threading.Event() + loop = asyncio.get_running_loop() + fenced: list[bool] = [] + loop.set_default_executor(ThreadPoolExecutor(max_workers=1)) + + @activity.defn(name="work") + def work() -> object: + context = activity.context() + loop.call_soon_threadsafe(entered.set) + try: + assert release.wait(timeout=2) + heartbeat = asyncio.run_coroutine_threadsafe(context.heartbeat(), loop) + try: + heartbeat.result(timeout=1) + except LocalActivityExecutionAborted: + fenced.append(True) + else: + fenced.append(False) + return object() + finally: + loop.call_soon_threadsafe(discarded.set) + + worker = await server.worker(monkeypatch, activities=[work], heartbeat_interval=0.01) + execution = worker._track(worker._run_workflow_task(claimed_task("local_activity", observed=False))) + await asyncio.wait_for(entered.wait(), timeout=1) + server.history.append(request()) + server.client.heartbeat_workflow_task.return_value = lease_ack(observed=True) + try: + commands = await asyncio.wait_for(execution, timeout=1) + assert commands is not None and [command["type"] for command in commands] == ["start_timer"] + heartbeats_after_delivery = server.client.heartbeat_workflow_task.await_count + finally: + release.set() + await asyncio.wait_for(discarded.wait(), timeout=1) + assert fenced == [True] + assert server.client.heartbeat_workflow_task.await_count == heartbeats_after_delivery + assert server.client.deliver_workflow_cancellation.await_args.kwargs["call_kind"] == "local_activity" + server.client.fail_workflow_task.assert_not_awaited() + await worker.stop() + + @pytest.mark.parametrize("value", ["", None, {}, "not-a-page"]) async def test_invalid_or_missing_refresh_token_never_delivers(monkeypatch: pytest.MonkeyPatch, value: Any) -> None: server = ClaimServer() From 119747bb2fb838f6dff64c9938c780414a074b99 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 02:23:21 +0000 Subject: [PATCH 06/41] Qualify cooperative cancellation against a connected candidate Server --- .github/workflows/ci.yml | 26 +- README.md | 13 +- docker-compose.test.yml | 2 + scripts/ci/checkout-public-repository.py | 20 +- tests/integration/cooperative_worker.py | 47 ++ .../test_cooperative_cancellation.py | 449 ++++++++++++++++++ tests/test_ci_checkout.py | 42 ++ 7 files changed, 594 insertions(+), 5 deletions(-) create mode 100644 tests/integration/cooperative_worker.py create mode 100644 tests/integration/test_cooperative_cancellation.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4683602..6d130f3 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,6 +6,15 @@ on: pull_request: branches: [main] workflow_dispatch: + inputs: + server_commit: + description: Exact public Server candidate SHA, or empty for its default branch + type: string + default: "" + cooperative_qualification: + description: Qualify the candidate cooperative protocol 1.20 + type: boolean + default: false permissions: contents: read @@ -145,6 +154,9 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 25 needs: [lint, test] + env: + DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION: ${{ inputs.cooperative_qualification && '1.20' || '1.19' }} + DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION: ${{ inputs.cooperative_qualification && '1' || '0' }} steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: @@ -153,7 +165,9 @@ jobs: with: python-version: "3.12" - name: Check out public Server integration source - run: python sdk-python/scripts/ci/checkout-public-repository.py server server + env: + SERVER_COMMIT: ${{ inputs.server_commit }} + run: python sdk-python/scripts/ci/checkout-public-repository.py server server --commit "$SERVER_COMMIT" - run: pip install -e '.[dev]' working-directory: sdk-python - name: Configure isolated Docker project @@ -171,7 +185,15 @@ jobs: working-directory: sdk-python env: DURABLE_WORKFLOW_AUTH_TOKEN: test-token - run: pytest tests/integration/ -v + run: pytest tests/integration/ -v --junitxml=integration-results.xml + - name: Retain connected scenario results + if: ${{ always() && github.server_url == 'https://github.com' }} + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7 + with: + name: connected-integration-${{ github.sha }} + path: sdk-python/integration-results.xml + if-no-files-found: warn + retention-days: 7 - name: Emit integration diagnostics if: failure() working-directory: sdk-python diff --git a/README.md b/README.md index 2ce8503..1641a15 100644 --- a/README.md +++ b/README.md @@ -127,11 +127,22 @@ pytest tests/ -m "not integration" Integration tests use Docker: ```bash +export COMPOSE_PROJECT_NAME=sdk-python-local docker compose -f docker-compose.test.yml up -d --build --wait -pytest tests/integration/ -v +SERVER_PORT=$(docker compose -f docker-compose.test.yml port server 8080 | sed 's/.*://') +DURABLE_WORKFLOW_SERVER_URL="http://127.0.0.1:$SERVER_PORT" DURABLE_WORKFLOW_AUTH_TOKEN=test-token pytest tests/integration/ -v docker compose -f docker-compose.test.yml down -v ``` +Candidate cooperative cancellation qualification is explicit. In a manual CI +run, supply an exact public `server_commit` and set `cooperative_qualification` +to true. CI verifies that checkout, builds the candidate Server, enables protocol +1.20, runs the connected cases and retains their JUnit results. For a local +candidate, set `DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION=1.20` before starting +Compose and `DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION=1` for pytest. These cases +fail if the runtime does not discover the required capability. Ordinary CI +keeps protocol 1.19 and skips this unpublished feature's connected cases. + ## License [MIT](LICENSE) diff --git a/docker-compose.test.yml b/docker-compose.test.yml index 9a1cc12..50d56bb 100644 --- a/docker-compose.test.yml +++ b/docker-compose.test.yml @@ -30,6 +30,8 @@ services: WORKFLOW_SERVER_AUTH_DRIVER: token WORKFLOW_SERVER_AUTH_TOKEN: "test-token" DW_WORKFLOW_TASK_TIMEOUT: 10 + # Candidate cooperative qualification is explicit. Published defaults remain 1.19. + DW_WORKER_PROTOCOL_VERSION: "${DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION:-1.19}" depends_on: mysql: condition: service_healthy diff --git a/scripts/ci/checkout-public-repository.py b/scripts/ci/checkout-public-repository.py index fe3c681..12758bf 100644 --- a/scripts/ci/checkout-public-repository.py +++ b/scripts/ci/checkout-public-repository.py @@ -5,6 +5,7 @@ import argparse import os +import re import subprocess from collections.abc import Sequence from pathlib import Path @@ -15,8 +16,10 @@ } -def checkout(repository: str, destination: Path) -> None: +def checkout(repository: str, destination: Path, commit: str = "") -> None: """Clone a supported public repository without runner-host credentials.""" + if commit and re.fullmatch(r"[0-9a-f]{40}", commit) is None: + raise ValueError("integration source commit must be a full lowercase Git SHA") environment = os.environ.copy() environment["GIT_TERMINAL_PROMPT"] = "0" @@ -34,18 +37,31 @@ def checkout(repository: str, destination: Path) -> None: check=True, env=environment, ) + if commit: + git = ["git", "-c", "credential.helper=", "-C", str(destination)] + subprocess.run([*git, "fetch", "--depth=1", "origin", commit], check=True, env=environment) + subprocess.run([*git, "checkout", "--detach", commit], check=True, env=environment) + resolved = subprocess.run( + [*git, "rev-parse", "HEAD"], check=True, env=environment, capture_output=True, text=True, + ).stdout.strip() + if resolved != commit: + raise RuntimeError("integration checkout did not resolve the requested commit") + print(f"Integration {repository} source commit: {resolved}") def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("repository", choices=sorted(PUBLIC_REPOSITORIES)) parser.add_argument("destination", type=Path) + parser.add_argument( + "--commit", default="", help="Exact public candidate SHA. Defaults to the repository default branch.", + ) return parser.parse_args(argv) def main(argv: Sequence[str] | None = None) -> int: args = parse_args(argv) - checkout(args.repository, args.destination) + checkout(args.repository, args.destination, args.commit) return 0 diff --git a/tests/integration/cooperative_worker.py b/tests/integration/cooperative_worker.py new file mode 100644 index 0000000..a60dabd --- /dev/null +++ b/tests/integration/cooperative_worker.py @@ -0,0 +1,47 @@ +"""A real child worker for cooperative process-loss qualification.""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys + +from durable_workflow import Client +from tests.integration.test_cooperative_cancellation import candidate_worker, cooperative_cleanup, poll_claim + + +async def main() -> None: + queue, worker_id, mode = sys.argv[1:] + if mode not in {"hold", "finish"}: + raise ValueError("unknown qualification mode") + async with Client( + os.environ["DURABLE_WORKFLOW_SERVER_URL"], + token=os.environ.get("DURABLE_WORKFLOW_AUTH_TOKEN", "test-token"), namespace="default", + ) as client: + worker = candidate_worker(client, queue, worker_id=worker_id) + + async def cleanup(request_id: str) -> str: + print(json.dumps({"phase": "cleanup", "request_id": request_id, "worker_id": worker_id}), flush=True) + if mode == "hold": + await asyncio.Event().wait() + return await cooperative_cleanup(request_id) + + worker.activities["tests.python-cooperative-cleanup"] = cleanup + await worker._register() + try: + task = await poll_claim(client, worker) + observation = task["cancellation_request"] + print(json.dumps({ + "phase": "claim", "attempt": task["workflow_task_attempt"], "worker_id": worker_id, + "task_id": task["task_id"], "request_id": observation["request_id"], + "cleanup_deadline_at": observation["cleanup_deadline_at"], + }), flush=True) + commands = await worker._run_workflow_task(task) + print(json.dumps({"phase": "finished", "committed": commands is not None}), flush=True) + finally: + await worker.stop() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py new file mode 100644 index 0000000..66cf002 --- /dev/null +++ b/tests/integration/test_cooperative_cancellation.py @@ -0,0 +1,449 @@ +"""Connected cooperative qualification, opt-in until the service tuple is published.""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys +import threading +import uuid +from typing import Any + +import pytest + +from durable_workflow import Client, Worker, activity, workflow +from durable_workflow.client import WorkflowHandle +from durable_workflow.errors import ServerError, WorkflowCancelled +from durable_workflow.workflow import LocalActivityExecutionAborted + +pytestmark = pytest.mark.usefixtures("cooperative_runtime") + + +@pytest.fixture +async def cooperative_runtime(server_url: str, server_token: str, monkeypatch: pytest.MonkeyPatch) -> None: + if os.environ.get("DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION") != "1": + pytest.skip("candidate cooperative Server qualification is opt-in") + async with Client(server_url, token=server_token, namespace="default") as client: + info = await client.get_cluster_info() + assert info["worker_protocol"]["server_capabilities"]["cooperative_cancellation"] is True + assert info["worker_protocol"]["version"] == "1.20" + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + + +@workflow.defn(name="tests.python-cooperative-cleanup") +class CooperativeCleanupWorkflow: + def run(self, ctx: Any, kind: str) -> Any: + try: + if kind == "local": + yield ctx.local_activity("tests.python-cooperative-work", []) + elif kind == "remote": + yield ctx.schedule_activity("tests.python-cooperative-work", []) + else: + yield ctx.start_timer(300) + except WorkflowCancelled as error: + with ctx.cancellation_shield(): + yield ctx.local_activity("tests.python-cooperative-cleanup", [error.request_id]) + return error.request_id + return "not cancelled" + + +@activity.defn(name="tests.python-cooperative-work") +async def cooperative_work() -> str: + return "ordinary result" + + +@activity.defn(name="tests.python-cooperative-cleanup") +async def cooperative_cleanup(request_id: str) -> str: + await activity.context().heartbeat({"request_id": request_id}) + return request_id + + +def candidate_worker(client: Client, queue: str, **kwargs: Any) -> Worker: + return Worker( + client, task_queue=queue, worker_id=kwargs.pop("worker_id", f"{queue}-worker"), + workflows=[CooperativeCleanupWorkflow], activities=[cooperative_work, cooperative_cleanup], + capabilities=["cooperative_cancellation"], **kwargs, + ) + + +async def poll_claim(client: Client, worker: Worker) -> dict[str, Any]: + async def poll() -> dict[str, Any]: + while True: + task = await client.poll_workflow_task( + worker_id=worker.worker_id, task_queue=worker.task_queue, timeout=worker._poll_http_timeout, + ) + if task is not None: + return task + await asyncio.sleep(0.1) + return await asyncio.wait_for(poll(), timeout=30) + + +async def events(handle: WorkflowHandle) -> list[dict[str, Any]]: + history = await handle.get_history(page_size=1000) + assert history.get("next_page_token") is None + return history.get("events", history.get("history_events", [])) + + +async def assert_cancelled_cleanup(handle: WorkflowHandle, request_id: str) -> list[dict[str, Any]]: + with pytest.raises(WorkflowCancelled): + await handle.result(timeout=10) + history = await events(handle) + kinds = [event["event_type"] for event in history] + cancellation_kinds = ("CooperativeCancellationRequested", "CooperativeCancellationDelivered", "WorkflowCancelled") + for kind in (*cancellation_kinds, "ActivityCompleted"): + assert kinds.count(kind) == 1 + for kind in ("WorkflowCompleted", "WorkflowFailed", "ActivityFailed", "ActivityTimedOut"): + assert kind not in kinds + for event in history: + if event["event_type"] in cancellation_kinds: + assert event["payload"]["workflow_command_id"] == request_id + return history + + +class LostDeliveryAcknowledgmentClient(Client): + delivery_calls = 0 + + async def deliver_workflow_cancellation(self, **kwargs: Any) -> dict[str, Any]: + result = await super().deliver_workflow_cancellation(**kwargs) + self.delivery_calls += 1 + if self.delivery_calls == 1: + raise ServerError(503, {"reason": "qualification_lost_delivery_acknowledgment"}) + return result + + +async def test_waiting_timer_is_cancelled_by_canonical_delivery( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-waiting-{uuid.uuid4().hex[:8]}" + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + task = await poll_claim(client, worker) + assert [command["type"] for command in await worker._run_workflow_task(task) or []] == ["start_timer"] + before = await events(handle) + assert [event["event_type"] for event in before].count("TimerScheduled") == 1 + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + resumed = await poll_claim(client, worker) + assert await worker._run_workflow_task(resumed) is not None + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + kinds = [event["event_type"] for event in history] + assert kinds.count("TimerCancelled") == 1 + assert "TimerFired" not in kinds + finally: + await worker.stop() + + +async def test_leased_remote_activity_cannot_complete_after_delivery( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-remote-{uuid.uuid4().hex[:8]}" + cleanup_entered = asyncio.Event() + release_cleanup = asyncio.Event() + + async def cleanup(request_id: str) -> str: + cleanup_entered.set() + await release_cleanup.wait() + return await cooperative_cleanup(request_id) + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + worker.activities["tests.python-cooperative-cleanup"] = cleanup + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + task = await poll_claim(client, worker) + assert [command["type"] for command in await worker._run_workflow_task(task) or []] == ["schedule_activity"] + remote = await client.poll_activity_task(worker_id=worker.worker_id, task_queue=queue, timeout=5) + assert remote is not None + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + resumed = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(resumed)) + await asyncio.wait_for(cleanup_entered.wait(), timeout=10) + late_fence = { + "task_id": remote["task_id"], "activity_attempt_id": remote["activity_attempt_id"], + "lease_owner": worker.worker_id, + } + heartbeat = await client.heartbeat_activity_task(**late_fence) + assert heartbeat["cancel_requested"] is True + assert heartbeat["can_continue"] is False + assert heartbeat["heartbeat_recorded"] is False + release_cleanup.set() + assert await asyncio.wait_for(execution, timeout=10) is not None + before = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + with pytest.raises(ServerError) as late_result: + await client.complete_activity_task(**late_fence, result="late remote result") + assert late_result.value.status == 409 + assert late_result.value.reason() == "run_cancelled" + with pytest.raises(ServerError) as late_failure: + await client.fail_activity_task( + **late_fence, message="qualification late failure", failure_type="qualification-late-failure", + ) + assert late_failure.value.status == 409 + assert late_failure.value.reason() == "run_cancelled" + assert await events(handle) == before + finally: + release_cleanup.set() + await worker.stop() + + +@pytest.mark.parametrize("lost_ack", [False, True]) +async def test_request_before_claim_retains_identity_and_proves_delivery( + server_url: str, server_token: str, lost_ack: bool, +) -> None: + queue = f"py-cooperative-request-{uuid.uuid4().hex[:8]}" + client_type = LostDeliveryAcknowledgmentClient if lost_ack else Client + async with client_type(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + repeated = await handle.request_cancellation(cleanup_timeout_seconds=300) + assert accepted["duplicate"] is False + assert repeated["duplicate"] is True + assert repeated["cancellation_request"] == accepted["cancellation_request"] + task = await poll_claim(client, worker) + assert task["cancellation_request"] == accepted["cancellation_request"] + commands = await worker._run_workflow_task(task) + assert commands is not None + assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"] + await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + if lost_ack: + assert isinstance(client, LostDeliveryAcknowledgmentClient) + assert client.delivery_calls == 1 + finally: + await worker.stop() + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +async def test_active_local_work_observes_request_and_discards_late_result( + server_url: str, server_token: str, handler_kind: str, +) -> None: + queue = f"py-cooperative-active-{uuid.uuid4().hex[:8]}" + entered = asyncio.Event() + discarded = asyncio.Event() + cleanup_entered = asyncio.Event() + release_thread = threading.Event() + loop = asyncio.get_running_loop() + + async def blocked_work() -> object: + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + discarded.set() + return object() + + def synchronous_work() -> object: + loop.call_soon_threadsafe(entered.set) + assert release_thread.wait(timeout=20) + loop.call_soon_threadsafe(discarded.set) + return object() + + async def cleanup(request_id: str) -> str: + cleanup_entered.set() + return await cooperative_cleanup(request_id) + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + worker.activities["tests.python-cooperative-work"] = ( + blocked_work if handler_kind == "async" else synchronous_work + ) + worker.activities["tests.python-cooperative-cleanup"] = cleanup + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["local"], + ) + task = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=10) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + await asyncio.wait_for(cleanup_entered.wait(), timeout=10) + release_thread.set() + await asyncio.wait_for(discarded.wait(), timeout=10) + assert await asyncio.wait_for(execution, timeout=10) is not None + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + completed = [event for event in history if event["event_type"] == "ActivityCompleted"] + assert completed[0]["payload"]["activity_type"] == "tests.python-cooperative-cleanup" + finally: + release_thread.set() + await worker.stop() + + +@pytest.mark.parametrize("shutdown_kind", ["drain", "timeout"]) +async def test_shutdown_during_shielded_cleanup_reuses_canonical_delivery( + server_url: str, server_token: str, shutdown_kind: str, +) -> None: + queue = f"py-cooperative-shutdown-{uuid.uuid4().hex[:8]}" + entered = asyncio.Event() + release = asyncio.Event() + late_fenced = asyncio.Event() + + async def cleanup(request_id: str) -> object: + entered.set() + try: + await release.wait() + except asyncio.CancelledError: + with pytest.raises(LocalActivityExecutionAborted): + await activity.context().heartbeat() + late_fenced.set() + return object() + return await cooperative_cleanup(request_id) + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, shutdown_timeout=5 if shutdown_kind == "drain" else 0.1) + worker.activities["tests.python-cooperative-cleanup"] = cleanup + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + task = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=10) + before = await events(handle) + stop = asyncio.create_task(worker.stop()) + if shutdown_kind == "drain": + await asyncio.sleep(0) + release.set() + await asyncio.wait_for(stop, timeout=10) + if shutdown_kind == "timeout": + await asyncio.wait_for(late_fenced.wait(), timeout=10) + assert execution.cancelled() or execution.result() is None + after = await events(handle) + assert after[:len(before)] == before + assert [event["event_type"] for event in after[len(before):]] == ["RepairRequested"] + assert after[-1]["payload"]["command"]["request_method"] == "DELETE" + assert after[-1]["payload"]["command"]["request_path"].endswith(worker.worker_id) + replacement = candidate_worker(client, queue, worker_id=f"{queue}-replacement") + await replacement._register() + try: + reclaimed = await poll_claim(client, replacement) + assert reclaimed["workflow_task_attempt"] > task["workflow_task_attempt"] + assert reclaimed["lease_owner"] == replacement.worker_id + original = accepted["cancellation_request"] + assert reclaimed["cancellation_request"]["request_id"] == original["request_id"] + assert reclaimed["cancellation_request"]["cleanup_deadline_at"] == original["cleanup_deadline_at"] + assert await replacement._run_workflow_task(reclaimed) is not None + finally: + await replacement.stop() + else: + assert await execution is not None + await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + finally: + release.set() + await worker.stop() + + +async def native_process(queue: str, worker_id: str, mode: str) -> asyncio.subprocess.Process: + return await asyncio.create_subprocess_exec( + sys.executable, "-m", "tests.integration.cooperative_worker", queue, worker_id, mode, + stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE, + ) + + +async def process_event(process: asyncio.subprocess.Process, phase: str) -> dict[str, Any]: + async def read() -> dict[str, Any]: + assert process.stdout is not None + while line := await process.stdout.readline(): + value = json.loads(line) + if value["phase"] == phase: + return value + assert process.stderr is not None + pytest.fail(f"worker exited before {phase}: {(await process.stderr.read()).decode()}") + return await asyncio.wait_for(read(), timeout=40) + + +async def test_killed_process_reclaims_cleanup_in_a_new_process( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-process-{uuid.uuid4().hex[:8]}" + processes: list[asyncio.subprocess.Process] = [] + async with Client(server_url, token=server_token, namespace="default") as client: + seed = candidate_worker(client, queue, worker_id=f"{queue}-seed") + await seed._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=120) + original = accepted["cancellation_request"] + first = await native_process(queue, f"{queue}-killed", "hold") + processes.append(first) + claim = await process_event(first, "claim") + cleanup = await process_event(first, "cleanup") + assert cleanup["request_id"] == original["request_id"] + before = await events(handle) + first.kill() + assert await asyncio.wait_for(first.wait(), timeout=10) == -9 + assert await events(handle) == before + replacement = await native_process(queue, f"{queue}-replacement", "finish") + processes.append(replacement) + reclaimed = await process_event(replacement, "claim") + assert reclaimed["attempt"] > claim["attempt"] + assert reclaimed["worker_id"] != claim["worker_id"] + assert reclaimed["request_id"] == original["request_id"] + assert reclaimed["cleanup_deadline_at"] == original["cleanup_deadline_at"] + resumed = await process_event(replacement, "cleanup") + assert resumed["request_id"] == original["request_id"] + assert (await process_event(replacement, "finished"))["committed"] is True + assert await asyncio.wait_for(replacement.wait(), timeout=10) == 0 + await assert_cancelled_cleanup(handle, original["request_id"]) + finally: + for process in processes: + if process.returncode is None: + process.kill() + await process.wait() + await seed.stop() + + +@pytest.mark.parametrize("close_kind", ["deadline", "terminate"]) +async def test_cleanup_deadline_and_termination_fence_in_flight_local_result( + server_url: str, server_token: str, close_kind: str, +) -> None: + queue = f"py-cooperative-close-{uuid.uuid4().hex[:8]}" + entered = asyncio.Event() + + async def cleanup(request_id: str) -> object: + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + return object() + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue) + worker.activities["tests.python-cooperative-cleanup"] = cleanup + await worker._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + ) + await handle.request_cancellation(cleanup_timeout_seconds=2 if close_kind == "deadline" else 60) + task = await poll_claim(client, worker) + execution = worker._track(worker._run_workflow_task(task)) + await asyncio.wait_for(entered.wait(), timeout=10) + if close_kind == "terminate": + await handle.terminate(reason="qualification termination during cleanup") + assert await asyncio.wait_for(execution, timeout=15) is None + history = await events(handle) + kinds = [event["event_type"] for event in history] + assert kinds.count("WorkflowCancelled" if close_kind == "deadline" else "WorkflowTerminated") == 1 + assert kinds.count("CooperativeCancellationRequested") == 1 + assert kinds.count("CooperativeCancellationDelivered") == 1 + assert "ActivityCompleted" not in kinds + assert "ActivityFailed" not in kinds + assert "WorkflowCompleted" not in kinds + finally: + await worker.stop() diff --git a/tests/test_ci_checkout.py b/tests/test_ci_checkout.py index 5ed297a..637bd27 100644 --- a/tests/test_ci_checkout.py +++ b/tests/test_ci_checkout.py @@ -72,3 +72,45 @@ def test_ci_workflow_uses_portable_public_checkouts() -> None: line.strip() for line in workflow.splitlines() if line.strip().startswith("repository: durable-workflow/") ] assert public_repository_inputs == ["repository: durable-workflow/.github"] + + +@pytest.mark.parametrize("commit", ["main", "short", "A" * 40, "--upload-pack=command", "0" * 40 + ";command"]) +def test_candidate_checkout_rejects_non_sha_before_running_git(tmp_path: Path, commit: str) -> None: + capture = tmp_path / "called" + fake_git = tmp_path / "git" + fake_git.write_text('#!/bin/sh\ntouch "$GIT_CAPTURE"\n') + fake_git.chmod(0o755) + environment = {**os.environ, "PATH": str(tmp_path), "GIT_CAPTURE": str(capture)} + result = subprocess.run( + [sys.executable, str(CHECKOUT_SCRIPT), "server", str(tmp_path / "server"), "--commit", commit], + env=environment, capture_output=True, text=True, + ) + assert result.returncode != 0 + assert not capture.exists() + + +@pytest.mark.parametrize("matches", [True, False]) +def test_candidate_checkout_verifies_the_requested_public_commit(tmp_path: Path, matches: bool) -> None: + commit = "a" * 40 + resolved = commit if matches else "b" * 40 + capture = tmp_path / "git-calls" + fake_git = tmp_path / "git" + fake_git.write_text( + '#!/bin/sh\nprintf "%s\\n" "$*" >> "$GIT_CAPTURE"\n' + 'case "$*" in *rev-parse*) printf "%s\\n" "$RESOLVED_COMMIT" ;; esac\n' + ) + fake_git.chmod(0o755) + environment = { + **os.environ, "PATH": str(tmp_path), "GIT_CAPTURE": str(capture), "RESOLVED_COMMIT": resolved, + } + result = subprocess.run( + [sys.executable, str(CHECKOUT_SCRIPT), "server", str(tmp_path / "server"), "--commit", commit], + env=environment, capture_output=True, text=True, + ) + assert (result.returncode == 0) is matches + calls = capture.read_text().splitlines() + assert "https://github.com/durable-workflow/server.git" in calls[0] + assert any(f"fetch --depth=1 origin {commit}" in call for call in calls) + assert all("credential.helper=" in call for call in calls) + if matches: + assert f"Integration server source commit: {commit}" in result.stdout From 8c6bd876effaad4e559f799f54ca3290cb430af8 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 06:29:46 +0000 Subject: [PATCH 07/41] Supervise cooperative remote activity callbacks with bounded ownership observation --- src/durable_workflow/client.py | 22 ++ src/durable_workflow/worker.py | 220 ++++++++++++++++- tests/integration/COOPERATIVE.md | 50 ++++ tests/integration/cooperative_worker.py | 15 +- .../test_cooperative_cancellation.py | 190 +++++++++++++++ tests/test_cooperative_remote_worker.py | 230 ++++++++++++++++++ 6 files changed, 714 insertions(+), 13 deletions(-) create mode 100644 tests/integration/COOPERATIVE.md create mode 100644 tests/test_cooperative_remote_worker.py diff --git a/src/durable_workflow/client.py b/src/durable_workflow/client.py index a46e481..d611e44 100644 --- a/src/durable_workflow/client.py +++ b/src/durable_workflow/client.py @@ -5553,6 +5553,28 @@ async def fail_activity_task( "POST", f"/worker/activity-tasks/{task_id}/fail", worker=True, json=body ) + async def activity_task_status( + self, + *, + task_id: str, + activity_attempt_id: str, + lease_owner: str, + ) -> Any: + """Observe one cooperative activity claim without recording progress. + + Requires explicit worker protocol 1.20. This readonly observation does + not renew the activity lease or its heartbeat deadline. A positive reply + is not a reservation of ownership for a subsequent completion. + """ + if not _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )): + raise ValueError("activity ownership observation requires explicit worker protocol 1.20") + return await asyncio.wait_for(self._request( + "POST", f"/worker/activity-tasks/{task_id}/status", worker=True, + json={"activity_attempt_id": activity_attempt_id, "lease_owner": lease_owner}, timeout=5.0, + ), timeout=5.0) + async def heartbeat_activity_task( self, *, diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 23fd226..66ef63e 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -29,7 +29,7 @@ import types import uuid from collections.abc import Awaitable, Callable, Coroutine, Iterable, Mapping -from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import Future, ThreadPoolExecutor from datetime import datetime, timezone from functools import wraps from types import FunctionType @@ -159,6 +159,10 @@ _R = TypeVar("_R") +class _RemoteActivityExecutionAborted(Exception): + """Ownership cannot authorize another remote callback boundary.""" + + class _LocalActivityTimedOut(Exception): def __init__(self, kind: str) -> None: super().__init__(f"local activity {kind} timeout elapsed") @@ -1060,6 +1064,9 @@ def __init__( self._stop = asyncio.Event() self._local_activity_shutdown = asyncio.Event() self._local_activity_executor: ThreadPoolExecutor | None = None + self._remote_activity_executor: ThreadPoolExecutor | None = None + self._remote_activity_threads: set[Future[Any]] = set() + self._remote_activity_thread_slots = asyncio.Semaphore(max_concurrent_activity_tasks) self._wf_semaphore = asyncio.Semaphore(max_concurrent_workflow_tasks) self._act_semaphore = asyncio.Semaphore(max_concurrent_activity_tasks) self._shutdown_timeout = shutdown_timeout @@ -2175,6 +2182,159 @@ async def _report_workflow_task_after_completion_error( fail_error, ) + async def _assert_remote_activity_claim(self, task: dict[str, Any]) -> None: + try: + if self._local_activity_shutdown.is_set(): + raise _RemoteActivityExecutionAborted("worker shutdown abandoned its remote activity claim") + reply = await self.client.activity_task_status( + task_id=task["task_id"], + activity_attempt_id=task.get("activity_attempt_id") or task.get("attempt_id", ""), + lease_owner=self.worker_id, + ) + if (not isinstance(reply, Mapping) + or reply.get("task_id") != task["task_id"] + or reply.get("activity_attempt_id") != ( + task.get("activity_attempt_id") or task.get("attempt_id", "")) + or reply.get("lease_owner") != self.worker_id + or reply.get("can_continue") is not True + or reply.get("cancel_requested") is not False + or reply.get("heartbeat_recorded") is not False + or reply.get("reason") is not None): + raise _RemoteActivityExecutionAborted("remote activity observation refused its ownership fence") + bounds = [reply.get("lease_expires_at")] + deadlines = reply.get("deadlines") + if deadlines is not None: + if not isinstance(deadlines, Mapping): + raise _RemoteActivityExecutionAborted("remote activity observation returned invalid deadlines") + bounds.extend(deadlines[kind] for kind in ("heartbeat", "start_to_close", "schedule_to_close") + if deadlines.get(kind) is not None) + session = reply.get("worker_session") + if session is not None: + if (not isinstance(session, Mapping) or session.get("status") != "active" + or session.get("lease_owner") != self.worker_id): + raise _RemoteActivityExecutionAborted("remote activity observation lost its worker session") + bounds.extend([session.get("lease_expires_at"), session.get("ttl_expires_at")]) + for bound in bounds: + if not isinstance(bound, str) or "T" not in bound: + raise _RemoteActivityExecutionAborted("remote activity observation returned an invalid deadline") + deadline = datetime.fromisoformat(bound.replace("Z", "+00:00")) + if deadline.tzinfo is None or deadline <= datetime.now(timezone.utc): + raise _RemoteActivityExecutionAborted("remote activity ownership or execution deadline elapsed") + if self._local_activity_shutdown.is_set(): + raise _RemoteActivityExecutionAborted("worker shutdown abandoned its remote activity claim") + except _RemoteActivityExecutionAborted: + raise + except Exception as error: + raise _RemoteActivityExecutionAborted("remote activity ownership observation failed") from error + + async def _execute_cooperative_remote_callable( + self, task: dict[str, Any], handler: Callable[..., Any], args: tuple[Any, ...], info: ActivityInfo, + ) -> Any: + abandoned = threading.Event() + owner_loop = asyncio.get_running_loop() + + def boundary() -> None: + if abandoned.is_set() or self._local_activity_shutdown.is_set(): + raise _RemoteActivityExecutionAborted("remote activity callback no longer owns its claim") + + async def observe() -> None: + boundary() + try: + await self._assert_remote_activity_claim(task) + except BaseException: + abandoned.set() + raise + boundary() + + async def send_heartbeat(details: dict[str, Any] | None) -> None: + await observe() + try: + reply = await asyncio.wait_for(self.client.heartbeat_activity_task( + task_id=info.task_id, activity_attempt_id=info.activity_attempt_id, + lease_owner=self.worker_id, details=details, + ), timeout=5.0) + if (not isinstance(reply, Mapping) or reply.get("task_id") != info.task_id + or reply.get("activity_attempt_id") != info.activity_attempt_id + or reply.get("lease_owner") != self.worker_id or reply.get("can_continue") is not True + or reply.get("cancel_requested") is not False): + raise _RemoteActivityExecutionAborted("remote activity heartbeat lost its ownership fence") + except Exception as error: + abandoned.set() + raise _RemoteActivityExecutionAborted("remote activity user heartbeat failed") from error + await observe() + + async def heartbeat(details: dict[str, Any] | None) -> None: + boundary() + if asyncio.get_running_loop() is owner_loop: + await send_heartbeat(details) + else: + # A synchronous handler may use its own loop to await heartbeat. + # Its HTTP client and claim checks remain on the owner loop. + proxy = asyncio.run_coroutine_threadsafe(send_heartbeat(details), owner_loop) + await asyncio.wrap_future(proxy) + + if inspect.iscoroutinefunction(handler): + async def guarded(*arguments: Any) -> Any: + boundary() + return await handler(*arguments) + callback: Callable[..., Any] = guarded + else: + def guarded_sync(*arguments: Any) -> Any: + boundary() + return handler(*arguments) + callback = guarded_sync + + async def invoke() -> Any: + _set_context(ActivityContext(info=info, client=self.client, heartbeat_callback=heartbeat)) + try: + return await self._execute_activity_callable( + task, info.activity_type, args, callback, run_sync_in_thread=True, remote=True, + ) + finally: + _set_context(None) + + async def observe_ownership() -> None: + while True: + await asyncio.sleep(1.0) + await observe() + + await observe() + invocation = asyncio.create_task(invoke()) + observation = asyncio.create_task(observe_ownership()) + shutdown = asyncio.create_task(self._local_activity_shutdown.wait()) + try: + done, _ = await asyncio.wait([invocation, observation, shutdown], return_when=asyncio.FIRST_COMPLETED) + if shutdown in done: + raise _RemoteActivityExecutionAborted("worker shutdown abandoned its remote activity claim") + if observation in done: + await observation + raise _RemoteActivityExecutionAborted("remote ownership observer stopped without a response") + try: + result = await invocation + except Exception: + await observe() # A genuine application failure also needs current authority. + raise + await observe() # Fence before the Client encodes or externalizes the result. + return result + finally: + abandoned.set() + shutdown.cancel() + observation.cancel() + with contextlib.suppress(asyncio.CancelledError, Exception): + await shutdown + with contextlib.suppress(asyncio.CancelledError, Exception): + await observation + if not invocation.done(): + invocation.cancel() + + def discard_late_result(future: asyncio.Task[Any]) -> None: + if not future.cancelled(): + future.exception() + + # A running thread or cancellation-resistant callable can outlive + # the await. Its heartbeat and eventual publication stay fenced. + invocation.add_done_callback(discard_late_result) + async def _run_activity_task(self, task: dict[str, Any]) -> str: self._track_worker_session_from_task(task) task_id: str = task["task_id"] @@ -2261,7 +2421,13 @@ async def _run_activity_task(self, task: dict[str, Any]) -> str: ) _set_context(act_ctx) try: - result = await self._execute_activity_callable(task, activity_type, tuple(args), fn) + if self._cooperative_cancellation_supported: + result = await self._execute_cooperative_remote_callable(task, fn, tuple(args), act_ctx.info) + else: + result = await self._execute_activity_callable(task, activity_type, tuple(args), fn) + except _RemoteActivityExecutionAborted as error: + log.warning("remote activity %s claim abandoned: %s", task_id, error) + return "claim_aborted" except ActivityCancelled: log.info("activity %s cancelled via heartbeat", task_id) try: @@ -2337,7 +2503,7 @@ async def _execute_activity_callable( activity_type: str, args: tuple[Any, ...], fn: Callable[..., Any], - *, run_sync_in_thread: bool = False, + *, run_sync_in_thread: bool = False, remote: bool = False, ) -> Any: context = ActivityInterceptorContext( worker_id=self.worker_id, @@ -2349,14 +2515,40 @@ async def _execute_activity_callable( async def call_activity(ctx: ActivityInterceptorContext) -> Any: if run_sync_in_thread and not inspect.iscoroutinefunction(fn): - if self._local_activity_executor is None: - # Replay threads wait for local results, so sharing their pool can deadlock. - self._local_activity_executor = ThreadPoolExecutor( - max_workers=self.max_concurrent_workflow_tasks, thread_name_prefix="dw-local-activity", + if remote: + await self._remote_activity_thread_slots.acquire() + if self._remote_activity_executor is None: + self._remote_activity_executor = ThreadPoolExecutor( + max_workers=self.max_concurrent_activity_tasks, thread_name_prefix="dw-remote-activity", + ) + owner_loop = asyncio.get_running_loop() + try: + future = self._remote_activity_executor.submit(contextvars.copy_context().run, fn, *ctx.args) + except BaseException: + self._remote_activity_thread_slots.release() + raise + self._remote_activity_threads.add(future) + + def thread_finished(completed: Future[Any]) -> None: + def release_capacity() -> None: + self._remote_activity_threads.discard(completed) + self._remote_activity_thread_slots.release() + with contextlib.suppress(RuntimeError): + owner_loop.call_soon_threadsafe(release_capacity) + + # Cancellation of the await cannot release a running thread's + # slot. Only actual completion or cancellation before start can. + future.add_done_callback(thread_finished) + result = await asyncio.wrap_future(future) + else: + if self._local_activity_executor is None: + # Replay threads wait for local results, so sharing their pool can deadlock. + self._local_activity_executor = ThreadPoolExecutor( + max_workers=self.max_concurrent_workflow_tasks, thread_name_prefix="dw-local-activity", + ) + result = await asyncio.get_running_loop().run_in_executor( + self._local_activity_executor, contextvars.copy_context().run, fn, *ctx.args, ) - result = await asyncio.get_running_loop().run_in_executor( - self._local_activity_executor, contextvars.copy_context().run, fn, *ctx.args, - ) else: result = fn(*ctx.args) if asyncio.iscoroutine(result): @@ -2808,6 +3000,9 @@ async def _report_unhandled_workflow_task_error( async def _poll_activity_tasks(self) -> None: while not self._stop.is_set(): + if len(self._remote_activity_threads) >= self.max_concurrent_activity_tasks: + await asyncio.sleep(0.1) + continue await self._act_semaphore.acquire() if self._stop.is_set(): self._act_semaphore.release() @@ -3242,7 +3437,8 @@ def _current_task_slots(self) -> dict[str, int]: 0, self.max_concurrent_workflow_tasks - self._workflow_reserved ), "activity_available": max( - 0, self.max_concurrent_activity_tasks - self._activity_inflight + 0, min(self.max_concurrent_activity_tasks - self._activity_inflight, + self.max_concurrent_activity_tasks - len(self._remote_activity_threads)), ), "session_available": max( 0, self.max_concurrent_worker_sessions @@ -3573,6 +3769,8 @@ async def _shutdown(self) -> None: if self._local_activity_executor is not None: self._local_activity_executor.shutdown(wait=False, cancel_futures=True) + if self._remote_activity_executor is not None: + self._remote_activity_executor.shutdown(wait=False, cancel_futures=True) for session in self._worker_sessions.values(): if not session.active: diff --git a/tests/integration/COOPERATIVE.md b/tests/integration/COOPERATIVE.md new file mode 100644 index 0000000..7923744 --- /dev/null +++ b/tests/integration/COOPERATIVE.md @@ -0,0 +1,50 @@ +# Cooperative cancellation source qualification + +The published SDK defaults to worker protocol 1.19. These cases explicitly select +the candidate 1.20 protocol and require compatible Server capability discovery. + +Run the `CI` workflow on the candidate branch with `cooperative_qualification=true` +and `server_commit` set to the exact 40-character public Server candidate SHA. +The workflow checks out that source, builds an isolated MySQL/Redis Server stack, +runs the integration suite, retains JUnit and removes its containers, images, +network and payload volume. Against an already isolated candidate stack: + +```sh +export DURABLE_WORKFLOW_SERVER_URL=http://127.0.0.1:8080 +export DURABLE_WORKFLOW_AUTH_TOKEN=test-token +export DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION=1.20 +export DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION=1 +pytest tests/integration/test_cooperative_cancellation.py -v --junitxml=cooperative-results.xml +``` + +The suite covers canonical delivery and cleanup, lost replies, waiting runs, +local callbacks, cleanup shutdown/replay and deadlines. Its remote cases use the +actual `Worker.run()` pollers and callback dispatcher, rather than only manually +claiming an attempt and sending a heartbeat. Async and synchronous callbacks are +tested with and without user heartbeats, through an actual accepted worker +registration heartbeat. Shutdown expiry fences callbacks before replacement +cleanup. Another case kills the real remote owner process, delivers the request +in a new process and rejects the dead owner's late completion/failure. + +Readonly status calls never record user progress or renew the activity lease. +They fail closed on refused/invalid ownership, elapsed execution/session bounds +or failed observation. A pending workflow request becomes an activity stop only +when the workflow worker durably delivers cancellation. Workflow capacity must +remain available while an activity is blocked. Here a Python worker provides it +through separate workflow execution and activity thread capacities. + +Synchronous remote handlers use a lazy pool bounded by configured activity +concurrency. A running thread retains its execution slot until it actually +finishes, even after its durable attempt is abandoned. Worker availability and +poll admission reflect that occupied capacity. Its late heartbeats and result +publication are fenced. Python cannot forcibly stop a running thread or a +callable that suppresses cancellation. Such code may continue external effects +until it returns or its process supervisor stops the process. Activities still +need idempotency and reconciliation. + +Only activity-authored heartbeats extend the existing five-minute activity +lease and any user heartbeat deadline. A positive readonly observation does not +reserve ownership for a later completion. The Server independently validates +completion and failure fences. Active remote-attempt reclaim under a new +activity owner and exact published Server/SDK qualification remain distinct +release gates. diff --git a/tests/integration/cooperative_worker.py b/tests/integration/cooperative_worker.py index a60dabd..e39162b 100644 --- a/tests/integration/cooperative_worker.py +++ b/tests/integration/cooperative_worker.py @@ -7,13 +7,13 @@ import os import sys -from durable_workflow import Client +from durable_workflow import Client, activity from tests.integration.test_cooperative_cancellation import candidate_worker, cooperative_cleanup, poll_claim async def main() -> None: queue, worker_id, mode = sys.argv[1:] - if mode not in {"hold", "finish"}: + if mode not in {"hold", "finish", "remote"}: raise ValueError("unknown qualification mode") async with Client( os.environ["DURABLE_WORKFLOW_SERVER_URL"], @@ -28,6 +28,17 @@ async def cleanup(request_id: str) -> str: return await cooperative_cleanup(request_id) worker.activities["tests.python-cooperative-cleanup"] = cleanup + if mode == "remote": + async def blocked_remote() -> object: + info = activity.context().info + print(json.dumps({"phase": "remote-entered", "task_id": info.task_id, + "activity_attempt_id": info.activity_attempt_id, + "lease_owner": info.worker_id}), flush=True) + await asyncio.Event().wait() + return object() + worker.activities["tests.python-cooperative-work"] = blocked_remote + await worker.run() + return await worker._register() try: task = await poll_claim(client, worker) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 66cf002..d965b35 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -15,6 +15,7 @@ from durable_workflow import Client, Worker, activity, workflow from durable_workflow.client import WorkflowHandle from durable_workflow.errors import ServerError, WorkflowCancelled +from durable_workflow.worker import _RemoteActivityExecutionAborted from durable_workflow.workflow import LocalActivityExecutionAborted pytestmark = pytest.mark.usefixtures("cooperative_runtime") @@ -112,6 +113,151 @@ async def deliver_workflow_cancellation(self, **kwargs: Any) -> dict[str, Any]: return result +class ObservedOwnerClient(Client): + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self.owner_heartbeat = asyncio.Event() + + async def heartbeat_worker(self, **kwargs: Any) -> Any: + reply = await super().heartbeat_worker(**kwargs) + if isinstance(reply, dict) and reply.get("acknowledged") is True: + self.owner_heartbeat.set() + return reply + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +@pytest.mark.parametrize("user_heartbeat", [False, True]) +async def test_actual_remote_worker_fences_blocked_callbacks_without_manufacturing_progress( + server_url: str, server_token: str, handler_kind: str, user_heartbeat: bool, +) -> None: + queue = f"py-cooperative-owner-{uuid.uuid4().hex[:8]}" + entered, late_fenced = asyncio.Event(), asyncio.Event() + release_thread = threading.Event() + loop = asyncio.get_running_loop() + fences: list[dict[str, str]] = [] + + def record_fence() -> None: + info = activity.context().info + fences.append({"task_id": info.task_id, "activity_attempt_id": info.activity_attempt_id, + "lease_owner": info.worker_id}) + + async def asynchronous() -> object: + record_fence() + entered.set() + try: + while True: + await asyncio.sleep(0.1) + if user_heartbeat: + await activity.context().heartbeat({"qualification": "remote-in-flight"}) + except (asyncio.CancelledError, _RemoteActivityExecutionAborted): + with pytest.raises(_RemoteActivityExecutionAborted): + await activity.context().heartbeat({"late": True}) + late_fenced.set() + return object() + + def synchronous() -> object: + record_fence() + loop.call_soon_threadsafe(entered.set) + try: + while not release_thread.wait(timeout=0.1): + if user_heartbeat: + asyncio.run(activity.context().heartbeat({"qualification": "remote-in-flight"})) + except _RemoteActivityExecutionAborted: + pass + with pytest.raises(_RemoteActivityExecutionAborted): + asyncio.run(activity.context().heartbeat({"late": True})) + loop.call_soon_threadsafe(late_fenced.set) + return object() + + async with ObservedOwnerClient(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, max_concurrent_activity_tasks=1) + worker.activities["tests.python-cooperative-work"] = asynchronous if handler_kind == "async" else synchronous + running = asyncio.create_task(worker.run()) + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + await asyncio.wait_for(entered.wait(), timeout=15) + await asyncio.wait_for(client.owner_heartbeat.wait(), timeout=15) + assert not late_fenced.is_set() + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + release_thread.set() + await asyncio.wait_for(late_fenced.wait(), timeout=5) + delivery = [event for event in history if event["event_type"] == "CooperativeCancellationDelivered"][0] + assert delivery["payload"]["call_kind"] == "activity" + assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 + progress = [event for event in history if event["event_type"] == "ActivityHeartbeatRecorded" + and event["payload"].get("activity_type") == "tests.python-cooperative-work"] + assert bool(progress) is user_heartbeat + assert len(fences) == 1 + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fences[0], result="late") + assert completion.value.status_code == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fences[0], message="late", failure_type="LateQualification") + assert failure.value.status_code == 409 + assert await events(handle) == history + finally: + release_thread.set() + await worker.stop() + await asyncio.wait_for(running, timeout=5) + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +async def test_actual_remote_shutdown_expiry_fences_callback_before_replacement_cleanup( + server_url: str, server_token: str, handler_kind: str, +) -> None: + queue = f"py-cooperative-owner-stop-{uuid.uuid4().hex[:8]}" + entered, late_fenced = asyncio.Event(), asyncio.Event() + release = threading.Event() + loop = asyncio.get_running_loop() + + async def asynchronous() -> object: + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + with pytest.raises(_RemoteActivityExecutionAborted): + await activity.context().heartbeat({"late": True}) + late_fenced.set() + return object() + + def synchronous() -> object: + loop.call_soon_threadsafe(entered.set) + assert release.wait(timeout=20) + with pytest.raises(_RemoteActivityExecutionAborted): + asyncio.run(activity.context().heartbeat({"late": True})) + loop.call_soon_threadsafe(late_fenced.set) + return object() + + async with Client(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, shutdown_timeout=0.1) + worker.activities["tests.python-cooperative-work"] = asynchronous if handler_kind == "async" else synchronous + running = asyncio.create_task(worker.run()) + replacement = candidate_worker(client, queue, worker_id=f"{queue}-replacement") + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + await asyncio.wait_for(entered.wait(), timeout=15) + before = await events(handle) + await asyncio.wait_for(worker.stop(), timeout=5) + await asyncio.wait_for(running, timeout=5) + release.set() + await asyncio.wait_for(late_fenced.wait(), timeout=5) + assert await events(handle) == before + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + await replacement._register() + assert await replacement._run_workflow_task(await poll_claim(client, replacement)) is not None + await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + finally: + release.set() + await replacement.stop() + await worker.stop() + await asyncio.wait_for(running, timeout=5) + + async def test_waiting_timer_is_cancelled_by_canonical_delivery( server_url: str, server_token: str, ) -> None: @@ -365,6 +511,50 @@ async def read() -> dict[str, Any]: return await asyncio.wait_for(read(), timeout=40) +async def test_killed_remote_owner_cannot_publish_after_cold_workflow_delivery( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-remote-process-{uuid.uuid4().hex[:8]}" + processes: list[asyncio.subprocess.Process] = [] + async with Client(server_url, token=server_token, namespace="default") as client: + seed = candidate_worker(client, queue, worker_id=f"{queue}-seed") + await seed._register() + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + owner = await native_process(queue, f"{queue}-killed", "remote") + processes.append(owner) + claimed = await process_event(owner, "remote-entered") + before = await events(handle) + owner.kill() + assert await asyncio.wait_for(owner.wait(), timeout=10) == -9 + assert await events(handle) == before + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + replacement = await native_process(queue, f"{queue}-replacement", "finish") + processes.append(replacement) + resumed = await process_event(replacement, "cleanup") + assert resumed["request_id"] == accepted["cancellation_request"]["request_id"] + assert (await process_event(replacement, "finished"))["committed"] is True + assert await asyncio.wait_for(replacement.wait(), timeout=10) == 0 + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 + fence = {key: claimed[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fence, result="late") + assert completion.value.status_code == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fence, message="late", failure_type="LateQualification") + assert failure.value.status_code == 409 + assert await events(handle) == history + finally: + for process in processes: + if process.returncode is None: + process.kill() + await process.wait() + await seed.stop() + + async def test_killed_process_reclaims_cleanup_in_a_new_process( server_url: str, server_token: str, ) -> None: diff --git a/tests/test_cooperative_remote_worker.py b/tests/test_cooperative_remote_worker.py new file mode 100644 index 0000000..1939112 --- /dev/null +++ b/tests/test_cooperative_remote_worker.py @@ -0,0 +1,230 @@ +from __future__ import annotations + +import asyncio +import threading +from copy import deepcopy +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +import pytest_asyncio + +from durable_workflow import activity, serializer +from durable_workflow.client import Client +from durable_workflow.errors import ActivityCancelled, NonRetryableError +from durable_workflow.retry_policy import TransportRetryPolicy +from durable_workflow.worker import Worker, _RemoteActivityExecutionAborted +from tests.test_worker import compatible_cluster_info + + +def task() -> dict[str, Any]: + return {"task_id": "remote-task", "activity_attempt_id": "attempt", "activity_type": "remote", + "payload_codec": "avro", "arguments": serializer.envelope([], codec="avro")} + + +def status() -> dict[str, Any]: + return {"task_id": "remote-task", "activity_attempt_id": "attempt", "lease_owner": "owner", + "can_continue": True, "cancel_requested": False, "reason": None, "heartbeat_recorded": False, + "lease_expires_at": "2100-01-01T00:00:00Z", "deadlines": None, "worker_session": None} + + +@pytest_asyncio.fixture +async def owner(monkeypatch: pytest.MonkeyPatch): # type: ignore[no-untyped-def] + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + client = AsyncMock(spec=Client) + client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": {"cooperative_cancellation": True}, + }) + client.register_worker.return_value = {"registered": True} + client.activity_task_status.return_value = status() + client.heartbeat_activity_task.return_value = status() + worker = Worker(client, task_queue="queue", worker_id="owner", capabilities=["cooperative_cancellation"], + max_concurrent_activity_tasks=1) + await worker._register() + try: + yield worker, client + finally: + await worker.stop() + + +@pytest.mark.parametrize("reason", ["cancel", "replacement", "backend", "shutdown"]) +async def test_blocked_async_callback_is_abandoned_without_progress_or_publication(owner, reason: str) -> None: + worker, client = owner + entered, late_fenced = asyncio.Event(), asyncio.Event() + + async def callback() -> object: + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + with pytest.raises(_RemoteActivityExecutionAborted): + await activity.context().heartbeat({"late": True}) + late_fenced.set() + return object() + + worker.activities["remote"] = callback + execution = asyncio.create_task(worker._run_activity_task(task())) + await asyncio.wait_for(entered.wait(), timeout=2) + if reason == "cancel": + client.activity_task_status.return_value = {**status(), "can_continue": False, + "cancel_requested": True, "reason": "activity_cancelled"} + elif reason == "replacement": + client.activity_task_status.return_value = {**status(), "activity_attempt_id": "replacement"} + elif reason == "backend": + client.activity_task_status.side_effect = OSError("backend unavailable") + else: + worker._local_activity_shutdown.set() + assert await asyncio.wait_for(execution, timeout=3) == "claim_aborted" + await asyncio.wait_for(late_fenced.wait(), timeout=2) + client.heartbeat_activity_task.assert_not_awaited() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +async def test_synchronous_callback_keeps_owner_available_and_fences_late_thread_heartbeat(owner) -> None: + worker, client = owner + entered, late_fenced = asyncio.Event(), asyncio.Event() + release = threading.Event() + loop = asyncio.get_running_loop() + + def callback() -> object: + loop.call_soon_threadsafe(entered.set) + assert release.wait(timeout=10) + with pytest.raises(_RemoteActivityExecutionAborted): + asyncio.run(activity.context().heartbeat({"late": True})) + loop.call_soon_threadsafe(late_fenced.set) + return object() + + worker.activities["remote"] = callback + execution = asyncio.create_task(worker._run_activity_task(task())) + try: + await asyncio.wait_for(entered.wait(), timeout=2) + client.activity_task_status.return_value = {**status(), "can_continue": False, "reason": "lease_expired"} + assert await asyncio.wait_for(execution, timeout=3) == "claim_aborted" + assert not late_fenced.is_set() # The Python thread is still running, with its attempt fenced. + assert worker._current_task_slots()["activity_available"] == 0 + poller = asyncio.create_task(worker._poll_activity_tasks()) + try: + await asyncio.sleep(0.05) + client.poll_activity_task.assert_not_awaited() + finally: + poller.cancel() + with pytest.raises(asyncio.CancelledError): + await poller + release.set() + await asyncio.wait_for(late_fenced.wait(), timeout=2) + client.heartbeat_activity_task.assert_not_awaited() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + finally: + release.set() + + +@pytest.mark.parametrize("synchronous", [False, True]) +async def test_authored_heartbeat_stays_on_owner_loop_and_preserves_typed_result(owner, synchronous: bool) -> None: + worker, client = owner + loop = asyncio.get_running_loop() + observed_loops = [] + + async def heartbeat(**kwargs: Any) -> dict[str, Any]: + observed_loops.append(asyncio.get_running_loop()) + assert kwargs["details"] == {"authored": True} + return status() + + async def asynchronous() -> dict[str, Any]: + await activity.context().heartbeat({"authored": True}) + return {"bytes": b"\x00\xff", "value": 42, "null": None} + + def blocking() -> dict[str, Any]: + asyncio.run(activity.context().heartbeat({"authored": True})) + return {"bytes": b"\x00\xff", "value": 42, "null": None} + + client.heartbeat_activity_task.side_effect = heartbeat + worker.activities["remote"] = blocking if synchronous else asynchronous + assert await worker._run_activity_task(task()) == "completed" + assert observed_loops == [loop] + assert client.complete_activity_task.await_args.kwargs["result"] == { + "bytes": b"\x00\xff", "value": 42, "null": None, + } + client.fail_activity_task.assert_not_awaited() + + +@pytest.mark.parametrize("field", ["lease", "heartbeat", "start_to_close", "schedule_to_close", "session", "malformed"]) +async def test_invalid_or_elapsed_observed_bounds_prevent_callback_start(owner, field: str) -> None: + worker, client = owner + reply = deepcopy(status()) + expired = "2000-01-01T00:00:00Z" + if field == "lease": + reply["lease_expires_at"] = expired + elif field == "malformed": + reply["deadlines"] = "invalid" + elif field == "session": + reply["worker_session"] = {"status": "active", "lease_owner": "owner", + "lease_expires_at": expired, "ttl_expires_at": "2100-01-01T00:00:00Z"} + else: + reply["deadlines"] = {field: expired} + client.activity_task_status.return_value = reply + callback = AsyncMock() + worker.activities["remote"] = callback + assert await worker._run_activity_task(task()) == "claim_aborted" + callback.assert_not_awaited() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +@pytest.mark.parametrize("error", [ValueError("original"), NonRetryableError("original"), ActivityCancelled()]) +async def test_genuine_application_failure_preserves_classification(owner, error: Exception) -> None: + worker, client = owner + + async def callback() -> None: + raise error + + worker.activities["remote"] = callback + outcome = await worker._run_activity_task(task()) + assert outcome in {"failed", "failed_non_retryable", "cancelled"} + assert client.fail_activity_task.await_args.kwargs["failure_type"] == type(error).__name__ + assert client.fail_activity_task.await_args.kwargs.get("non_retryable", False) is isinstance( + error, (NonRetryableError, ActivityCancelled), + ) + client.complete_activity_task.assert_not_awaited() + + +async def test_client_observation_keeps_worker_credentials_namespace_and_fence(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + async with Client("http://runtime.test", control_token="control", worker_token="worker", namespace="ns") as client: + reply = httpx.Response(200, json=status(), request=httpx.Request("POST", "http://runtime.test")) + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=reply) as send: + assert await client.activity_task_status( + task_id="remote-task", activity_attempt_id="attempt", lease_owner="owner", + ) == status() + headers = send.await_args.kwargs["headers"] + assert headers["Authorization"] == "Bearer worker" + assert headers["X-Namespace"] == "ns" + assert headers["X-Durable-Workflow-Protocol-Version"] == "1.20" + assert send.await_args.kwargs["json"] == {"activity_attempt_id": "attempt", "lease_owner": "owner"} + + +async def test_client_observation_requires_explicit_protocol_before_io(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + async with Client("http://runtime.test") as client: + with (patch.object(client._http, "request", new_callable=AsyncMock) as send, + pytest.raises(ValueError, match="1.20")): + await client.activity_task_status(task_id="task", activity_attempt_id="attempt", lease_owner="owner") + send.assert_not_awaited() + + +async def test_client_observation_bounds_the_total_retry_budget(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + retry = TransportRetryPolicy(initial_backoff_seconds=100, max_backoff_seconds=100, jitter=False) + async with Client("http://runtime.test", retry_policy=retry) as client: + refused = httpx.Response(503, json={"reason": "backend_unavailable"}, + request=httpx.Request("POST", "http://runtime.test")) + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=refused) as send: + started = asyncio.get_running_loop().time() + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(client.activity_task_status( + task_id="task", activity_attempt_id="attempt", lease_owner="owner", + ), timeout=7) + assert asyncio.get_running_loop().time() - started < 6 + send.assert_awaited_once() From 9f367496d03e5f17c5ee3bbd80a4bff1e1c25afd Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 06:31:17 +0000 Subject: [PATCH 08/41] Use the SDK ServerError status in connected publication fence assertions --- tests/integration/test_cooperative_cancellation.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index d965b35..98181ba 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -193,10 +193,10 @@ def synchronous() -> object: assert len(fences) == 1 with pytest.raises(ServerError) as completion: await client.complete_activity_task(**fences[0], result="late") - assert completion.value.status_code == 409 + assert completion.value.status == 409 with pytest.raises(ServerError) as failure: await client.fail_activity_task(**fences[0], message="late", failure_type="LateQualification") - assert failure.value.status_code == 409 + assert failure.value.status == 409 assert await events(handle) == history finally: release_thread.set() @@ -542,10 +542,10 @@ async def test_killed_remote_owner_cannot_publish_after_cold_workflow_delivery( fence = {key: claimed[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} with pytest.raises(ServerError) as completion: await client.complete_activity_task(**fence, result="late") - assert completion.value.status_code == 409 + assert completion.value.status == 409 with pytest.raises(ServerError) as failure: await client.fail_activity_task(**fence, message="late", failure_type="LateQualification") - assert failure.value.status_code == 409 + assert failure.value.status == 409 assert await events(handle) == history finally: for process in processes: From 88f0e31bf1736271deaabcc67d74df9ce98df491 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 06:35:12 +0000 Subject: [PATCH 09/41] Finish fenced remote claims without delaying worker shutdown on observer cancellation --- src/durable_workflow/worker.py | 23 +++++++++------------- tests/test_cooperative_remote_worker.py | 26 +++++++++++++++++++++++++ 2 files changed, 35 insertions(+), 14 deletions(-) diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 66ef63e..f0399a5 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -2318,22 +2318,17 @@ async def observe_ownership() -> None: return result finally: abandoned.set() - shutdown.cancel() - observation.cancel() - with contextlib.suppress(asyncio.CancelledError, Exception): - await shutdown - with contextlib.suppress(asyncio.CancelledError, Exception): - await observation - if not invocation.done(): - invocation.cancel() - def discard_late_result(future: asyncio.Task[Any]) -> None: - if not future.cancelled(): - future.exception() + def discard_late_result(future: asyncio.Task[Any]) -> None: + if not future.cancelled(): + future.exception() - # A running thread or cancellation-resistant callable can outlive - # the await. Its heartbeat and eventual publication stay fenced. - invocation.add_done_callback(discard_late_result) + # The abandoned claim must finish without waiting for callback or + # observer cancellation. All late progress/publication stays fenced. + for background in (invocation, observation, shutdown): + if not background.done(): + background.cancel() + background.add_done_callback(discard_late_result) async def _run_activity_task(self, task: dict[str, Any]) -> str: self._track_worker_session_from_task(task) diff --git a/tests/test_cooperative_remote_worker.py b/tests/test_cooperative_remote_worker.py index 1939112..c9efe27 100644 --- a/tests/test_cooperative_remote_worker.py +++ b/tests/test_cooperative_remote_worker.py @@ -121,6 +121,32 @@ def callback() -> object: release.set() +async def test_shutdown_expiry_of_tracked_remote_attempt_fences_a_cancellation_resistant_result(owner) -> None: + worker, client = owner + worker._shutdown_timeout = 0.01 + entered, late_fenced = asyncio.Event(), asyncio.Event() + + async def callback() -> object: + entered.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + with pytest.raises(_RemoteActivityExecutionAborted): + await activity.context().heartbeat() + late_fenced.set() + return object() + + worker.activities["remote"] = callback + execution = worker._track(worker._run_activity_task(task())) + await asyncio.wait_for(entered.wait(), timeout=2) + await asyncio.wait_for(worker.stop(), timeout=2) + await asyncio.wait_for(late_fenced.wait(), timeout=2) + assert execution.cancelled() or execution.result() == "claim_aborted" + client.heartbeat_activity_task.assert_not_awaited() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + @pytest.mark.parametrize("synchronous", [False, True]) async def test_authored_heartbeat_stays_on_owner_loop_and_preserves_typed_result(owner, synchronous: bool) -> None: worker, client = owner From ec794e2a1be5a89c7f282c580a255741f1b05be5 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 18:55:06 +0000 Subject: [PATCH 10/41] Qualify real Python activity process death and attempt reclaim --- .github/workflows/ci.yml | 2 +- docker-compose.test.yml | 9 +++ tests/integration/COOPERATIVE.md | 10 ++- tests/integration/cooperative_worker.py | 3 +- .../test_cooperative_cancellation.py | 69 ++++++++++++++++++- 5 files changed, 86 insertions(+), 7 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 6d130f3..cea83cb 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -198,7 +198,7 @@ jobs: if: failure() working-directory: sdk-python run: | - for service in bootstrap server worker; do + for service in bootstrap server worker repair; do echo "::group::$service logs" docker compose --project-name "$COMPOSE_PROJECT_NAME" -f docker-compose.test.yml \ logs --no-color --tail 200 "$service" || true diff --git a/docker-compose.test.yml b/docker-compose.test.yml index 50d56bb..9382016 100644 --- a/docker-compose.test.yml +++ b/docker-compose.test.yml @@ -72,6 +72,15 @@ services: redis: condition: service_healthy + repair: + build: + context: ../server + command: php artisan workflow:v2:repair-pass --loop --sleep-seconds=1 --json + environment: *server-env + depends_on: + server: + condition: service_healthy + mysql: image: mysql:8.0 environment: diff --git a/tests/integration/COOPERATIVE.md b/tests/integration/COOPERATIVE.md index 7923744..17e71d2 100644 --- a/tests/integration/COOPERATIVE.md +++ b/tests/integration/COOPERATIVE.md @@ -25,6 +25,11 @@ tested with and without user heartbeats, through an actual accepted worker registration heartbeat. Shutdown expiry fences callbacks before replacement cleanup. Another case kills the real remote owner process, delivers the request in a new process and rejects the dead owner's late completion/failure. +A separate case kills an active owner, waits for the real five-minute activity +lease and repair pass, and requires attempt 2 under a distinct owner before +requesting cancellation. It checks the killed attempt cannot publish a result, +failure or heartbeat, renew its lease, or change canonical history, then verifies +one cleanup with the original request identity. Readonly status calls never record user progress or renew the activity lease. They fail closed on refused/invalid ownership, elapsed execution/session bounds @@ -45,6 +50,5 @@ need idempotency and reconciliation. Only activity-authored heartbeats extend the existing five-minute activity lease and any user heartbeat deadline. A positive readonly observation does not reserve ownership for a later completion. The Server independently validates -completion and failure fences. Active remote-attempt reclaim under a new -activity owner and exact published Server/SDK qualification remain distinct -release gates. +completion and failure fences. Exact published Server/SDK qualification remains +a separate release gate from these source scenarios. diff --git a/tests/integration/cooperative_worker.py b/tests/integration/cooperative_worker.py index e39162b..3b95e26 100644 --- a/tests/integration/cooperative_worker.py +++ b/tests/integration/cooperative_worker.py @@ -33,7 +33,8 @@ async def blocked_remote() -> object: info = activity.context().info print(json.dumps({"phase": "remote-entered", "task_id": info.task_id, "activity_attempt_id": info.activity_attempt_id, - "lease_owner": info.worker_id}), flush=True) + "lease_owner": info.worker_id, + "attempt_number": info.attempt_number}), flush=True) await asyncio.Event().wait() return object() worker.activities["tests.python-cooperative-work"] = blocked_remote diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 98181ba..6a3fe29 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -7,7 +7,9 @@ import os import sys import threading +import time import uuid +from datetime import datetime from typing import Any import pytest @@ -499,7 +501,7 @@ async def native_process(queue: str, worker_id: str, mode: str) -> asyncio.subpr ) -async def process_event(process: asyncio.subprocess.Process, phase: str) -> dict[str, Any]: +async def process_event(process: asyncio.subprocess.Process, phase: str, timeout: float = 40) -> dict[str, Any]: async def read() -> dict[str, Any]: assert process.stdout is not None while line := await process.stdout.readline(): @@ -508,7 +510,70 @@ async def read() -> dict[str, Any]: return value assert process.stderr is not None pytest.fail(f"worker exited before {phase}: {(await process.stderr.read()).decode()}") - return await asyncio.wait_for(read(), timeout=40) + return await asyncio.wait_for(read(), timeout=timeout) + + +async def test_sigkill_activity_owner_reclaims_attempt_before_cooperative_cleanup( + server_url: str, server_token: str, +) -> None: + queue = f"py-cooperative-activity-reclaim-{uuid.uuid4().hex[:8]}" + processes: list[asyncio.subprocess.Process] = [] + async with Client(server_url, token=server_token, namespace="default") as client: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + ) + try: + owner = await native_process(queue, f"{queue}-killed", "remote") + processes.append(owner) + original = await process_event(owner, "remote-entered") + assert original["attempt_number"] == 1 + fence = {key: original[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} + leased = await client.activity_task_status(**fence) + assert leased["can_continue"] is True + assert leased["attempt_status"] == "running" + expires_at = datetime.fromisoformat(leased["lease_expires_at"].replace("Z", "+00:00")).timestamp() + assert expires_at > time.time(), "kill an actually current lease" + print(f"SIGKILL original activity: {json.dumps([original, leased])}") + owner.kill() + assert await asyncio.wait_for(owner.wait(), timeout=10) == -9 + + successor = await native_process(queue, f"{queue}-successor", "remote") + processes.append(successor) + # Wait for Native's actual five-minute lease and normal repair. + # No clock, storage row or production lease is changed. + reclaimed = await process_event(successor, "remote-entered", timeout=330) + assert reclaimed["task_id"] == original["task_id"] + assert reclaimed["activity_attempt_id"] != original["activity_attempt_id"] + assert reclaimed["lease_owner"] != original["lease_owner"] + assert reclaimed["attempt_number"] == 2 + assert time.time() >= expires_at, "reclaim precedes actual lease expiry" + print(f"SIGKILL successor activity: {json.dumps(reclaimed)}") + + before = await events(handle) + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fence, result="late") + assert completion.value.status == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fence, message="late", failure_type="LateQualification") + assert failure.value.status == 409 + heartbeat = await client.heartbeat_activity_task(**fence) + assert heartbeat["can_continue"] is False + assert heartbeat["heartbeat_recorded"] is False + assert heartbeat["cancel_requested"] is False + assert heartbeat["reason"] == "attempt_closed" + assert heartbeat["lease_expires_at"] == leased["lease_expires_at"] + assert heartbeat["last_heartbeat_at"] is None + assert await events(handle) == before, "dead attempt changed canonical history" + + accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 + print(f"SIGKILL reclaimed cancellation history: {json.dumps(history)}") + finally: + for process in processes: + if process.returncode is None: + process.kill() + await asyncio.wait_for(process.wait(), timeout=10) async def test_killed_remote_owner_cannot_publish_after_cold_workflow_delivery( From b3c45b67db136b1505ba3d79b8a38eaff6afa3b5 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 18:58:57 +0000 Subject: [PATCH 11/41] Prove expired attempt remains unchanged after heartbeat --- tests/integration/test_cooperative_cancellation.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 6a3fe29..e1dfc5d 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -549,6 +549,9 @@ async def test_sigkill_activity_owner_reclaims_attempt_before_cooperative_cleanu assert time.time() >= expires_at, "reclaim precedes actual lease expiry" print(f"SIGKILL successor activity: {json.dumps(reclaimed)}") + closed = await client.activity_task_status(**fence) + assert closed["attempt_status"] == "expired" + assert closed["can_continue"] is False before = await events(handle) with pytest.raises(ServerError) as completion: await client.complete_activity_task(**fence, result="late") @@ -561,8 +564,9 @@ async def test_sigkill_activity_owner_reclaims_attempt_before_cooperative_cleanu assert heartbeat["heartbeat_recorded"] is False assert heartbeat["cancel_requested"] is False assert heartbeat["reason"] == "attempt_closed" - assert heartbeat["lease_expires_at"] == leased["lease_expires_at"] - assert heartbeat["last_heartbeat_at"] is None + assert heartbeat["lease_expires_at"] == closed["lease_expires_at"] + assert heartbeat["last_heartbeat_at"] == closed["last_heartbeat_at"] + assert await client.activity_task_status(**fence) == closed assert await events(handle) == before, "dead attempt changed canonical history" accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) From e69cf8559aef46e2c069ed36882bd6f50da8105c Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 19:10:08 +0000 Subject: [PATCH 12/41] Use Worker backpressure policy in connected claim fixture --- .../integration/test_cooperative_cancellation.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index e1dfc5d..457d452 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -17,7 +17,7 @@ from durable_workflow import Client, Worker, activity, workflow from durable_workflow.client import WorkflowHandle from durable_workflow.errors import ServerError, WorkflowCancelled -from durable_workflow.worker import _RemoteActivityExecutionAborted +from durable_workflow.worker import _RemoteActivityExecutionAborted, _poll_capacity_delay from durable_workflow.workflow import LocalActivityExecutionAborted pytestmark = pytest.mark.usefixtures("cooperative_runtime") @@ -73,9 +73,17 @@ def candidate_worker(client: Client, queue: str, **kwargs: Any) -> Worker: async def poll_claim(client: Client, worker: Worker) -> dict[str, Any]: async def poll() -> dict[str, Any]: while True: - task = await client.poll_workflow_task( - worker_id=worker.worker_id, task_queue=worker.task_queue, timeout=worker._poll_http_timeout, - ) + try: + task = await client.poll_workflow_task( + worker_id=worker.worker_id, task_queue=worker.task_queue, timeout=worker._poll_http_timeout, + ) + except ServerError as error: + delay = _poll_capacity_delay(error, "workflow_task", worker.task_queue) + if delay is None: + raise + print(json.dumps({"phase": "poll-deferral", "delay_seconds": delay}), flush=True) + await asyncio.sleep(delay) + continue if task is not None: return task await asyncio.sleep(0.1) From 07e20300d5ba8d537a02e503074dc342a5b0823e Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 19:17:22 +0000 Subject: [PATCH 13/41] test: align repair cadence and fixture imports --- docker-compose.test.yml | 3 ++- tests/integration/test_cooperative_cancellation.py | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/docker-compose.test.yml b/docker-compose.test.yml index 9382016..c82ed76 100644 --- a/docker-compose.test.yml +++ b/docker-compose.test.yml @@ -75,7 +75,8 @@ services: repair: build: context: ../server - command: php artisan workflow:v2:repair-pass --loop --sleep-seconds=1 --json + # Match the production scheduler's unscoped, unthrottled repair cadence. + command: sh -c 'while true; do php artisan workflow:v2:repair-pass --json; sleep 10; done' environment: *server-env depends_on: server: diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 457d452..6a10040 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -17,7 +17,7 @@ from durable_workflow import Client, Worker, activity, workflow from durable_workflow.client import WorkflowHandle from durable_workflow.errors import ServerError, WorkflowCancelled -from durable_workflow.worker import _RemoteActivityExecutionAborted, _poll_capacity_delay +from durable_workflow.worker import _poll_capacity_delay, _RemoteActivityExecutionAborted from durable_workflow.workflow import LocalActivityExecutionAborted pytestmark = pytest.mark.usefixtures("cooperative_runtime") From 2b28ccad73b7d2ad76d9e5a4dad44a0bd33f2769 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 21:43:19 +0000 Subject: [PATCH 14/41] fix: defer Python parent claims while child cleanup waits --- src/durable_workflow/client.py | 17 ++++- src/durable_workflow/worker.py | 13 +++- tests/test_cooperative_cancellation_client.py | 64 +++++++++++++++++++ tests/test_cooperative_cancellation_worker.py | 49 +++++++++++++- 4 files changed, 140 insertions(+), 3 deletions(-) diff --git a/src/durable_workflow/client.py b/src/durable_workflow/client.py index d611e44..311a087 100644 --- a/src/durable_workflow/client.py +++ b/src/durable_workflow/client.py @@ -5163,7 +5163,9 @@ async def deliver_workflow_cancellation( Requires worker protocol 1.20 and a capable recorded claim. Retry with the same request, attempt and operation range after a lost acknowledgment, then reload canonical history before cleanup replay. - This method does not release or renew the lease. + A child wait may explicitly release the claim. Return to polling and + replay the boundary on the successor claim when the child finishes. + Successful delivery does not release or renew the lease. """ version = _protocol_version_from_env("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION) if not _supports_cooperative_cancellation_protocol(version): @@ -5194,6 +5196,19 @@ async def deliver_workflow_cancellation( "POST", f"/worker/workflow-tasks/{quote(task_id, safe='._:-')}/deliver-cancellation", worker=True, json=body, ) + if ( + isinstance(result, dict) + and result.get("delivered") is False + and delivery.call_kind in {"child", "parallel", "selection_handle"} + and result.get("reason") == "cancellation_waiting_for_child" + and result.get("claim_released") is True + and result.get("task_id") == task_id + and all(result.get(field) is None for field in ( + "request_id", "sequence", "call_kind", "sequence_span", + "operation_sequence", "operation_sequence_span", + )) + ): + return result if not isinstance(result, dict) or result.get("delivered") is not True or result.get("task_id") != task_id: raise ServerError(200, {"reason": "invalid_cooperative_cancellation_delivery"}) try: diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index f0399a5..4a74cd8 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -173,6 +173,10 @@ class _CooperativeCancellationObserved(LocalActivityExecutionAborted): """Return transport observation to the worker, never to authored cleanup.""" +class _WorkflowClaimDeferred(LocalActivityExecutionAborted): + """Server parked the parent and released its claim until child cleanup ends.""" + + class _InvalidLocalActivityReport(NonRetryableError): pass @@ -1523,13 +1527,17 @@ async def _replay_workflow_claim( observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) delivery_error: Exception | None = None try: - await self.client.deliver_workflow_cancellation( + delivery_reply = await self.client.deliver_workflow_cancellation( task_id=task["task_id"], lease_owner=self.worker_id, workflow_task_attempt=task.get("workflow_task_attempt", 1), request_id=intent.request_id, sequence=intent.sequence, call_kind=intent.call_kind, sequence_span=intent.sequence_span, operation_sequence=intent.operation_sequence, operation_sequence_span=intent.operation_sequence_span, ) + if delivery_reply.get("delivered") is False: + raise _WorkflowClaimDeferred("parent awaits canonical child cleanup on a new claim") + except _WorkflowClaimDeferred: + raise except Exception as error: delivery_error = error history = await self._refresh_cancellation_history(task, observed) @@ -1967,6 +1975,9 @@ def execute_local(command: RecordLocalActivity) -> Any: outcome, history = await self._replay_workflow_claim( cls, task, history, start_input, payload_codec=codec, execute_local=execute_local, ) + except _WorkflowClaimDeferred: + log.info("workflow task %s parked until child cleanup finishes", task_id) + return None except LocalActivityExecutionAborted as e: log.warning("abandoning workflow task %s before local activity commit: %s", task_id, e) return None diff --git a/tests/test_cooperative_cancellation_client.py b/tests/test_cooperative_cancellation_client.py index 8364260..476002e 100644 --- a/tests/test_cooperative_cancellation_client.py +++ b/tests/test_cooperative_cancellation_client.py @@ -36,6 +36,15 @@ def response(body: Any, status: int = 200) -> httpx.Response: return httpx.Response(status, json=body, request=httpx.Request("POST", "http://test")) +def pending_delivery_response(**overrides: Any) -> dict[str, Any]: + return { + "delivered": False, "task_id": "task/1", "workflow_run_id": "run-1", + "reason": "cancellation_waiting_for_child", "claim_released": True, + "request_id": None, "sequence": None, "call_kind": None, "sequence_span": None, + "operation_sequence": None, "operation_sequence_span": None, **overrides, + } + + @pytest_asyncio.fixture async def client() -> AsyncIterator[Client]: async with Client( @@ -180,6 +189,61 @@ async def test_delivery_sends_owner_attempt_and_authored_boundary( } +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["child", "parallel", "selection_handle"]) +async def test_pending_child_requires_explicit_claim_release( + client: Client, monkeypatch: pytest.MonkeyPatch, kind: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options = {"operation_sequence": 1} if kind == "selection_handle" else {} + pending = pending_delivery_response() + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=response(pending)): + result = await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind=kind, **options, + ) + assert result == pending + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"claim_released": False}, {"claim_released": "true"}, {"claim_released": None}, + {"task_id": "other"}, {"reason": "other"}, {"request_id": "original-request"}, + {"sequence": 3}, {"call_kind": "child"}, {"sequence_span": 1}, + {"operation_sequence": 1}, {"operation_sequence_span": 1}, {"delivered": 0}, +]) +async def test_malformed_pending_child_ack_is_rejected( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict[str, Any], +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(pending_delivery_response(**change))), + pytest.raises(ServerError) as error, + ): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="child", + ) + assert error.value.reason() == "invalid_cooperative_cancellation_delivery" + + +@pytest.mark.asyncio +async def test_pending_child_reply_cannot_release_an_unrelated_timer_claim( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(pending_delivery_response())), + pytest.raises(ServerError), + ): + await client.deliver_workflow_cancellation( + task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, + request_id="original-request", sequence=3, call_kind="timer", + ) + + @pytest.mark.asyncio @pytest.mark.parametrize("version", ["1.19", "2.20", "malformed"]) async def test_delivery_requires_explicit_compatible_protocol( diff --git a/tests/test_cooperative_cancellation_worker.py b/tests/test_cooperative_cancellation_worker.py index 4c0cb46..fb978fa 100644 --- a/tests/test_cooperative_cancellation_worker.py +++ b/tests/test_cooperative_cancellation_worker.py @@ -27,6 +27,8 @@ def run(self, ctx: workflow.WorkflowContext, kind: str): # type: ignore[no-unty yield ctx.local_activity("work", []) elif kind == "activity": yield ctx.schedule_activity("work", []) + elif kind == "child": + yield ctx.start_child_workflow("child", []) elif kind == "parallel": yield [ctx.start_timer(20), ctx.schedule_activity("work", [])] else: @@ -76,14 +78,17 @@ class ClaimServer: def __init__( self, *, history: list[dict[str, Any]] | None = None, lease_owner: str = "cooperative-worker", workflow_task_attempt: int = 4, + task_id: str = "task-1", ) -> None: self.client = AsyncMock(spec=Client) self.lease_owner = lease_owner self.workflow_task_attempt = workflow_task_attempt + self.task_id = task_id self.history = list(history if history is not None else [request()]) self.trace: list[str] = [] self.delivery_error: Exception | None = None self.commit_delivery = True + self.pending_deliveries = 0 self.client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ "version": "1.20", "server_capabilities": { "query_tasks": True, "long_poll_timeout": 30, "cooperative_cancellation": True, @@ -100,7 +105,7 @@ def __init__( async def page(self, **kwargs: Any) -> dict[str, Any]: assert kwargs == { - "task_id": "task-1", "next_history_page_token": "opaque-first-page", + "task_id": self.task_id, "next_history_page_token": "opaque-first-page", "lease_owner": self.lease_owner, "workflow_task_attempt": self.workflow_task_attempt, } self.trace.append("history") @@ -108,6 +113,12 @@ async def page(self, **kwargs: Any) -> dict[str, Any]: async def deliver(self, **kwargs: Any) -> dict[str, Any]: self.trace.append("delivery") + if self.pending_deliveries: + self.pending_deliveries -= 1 + return { + "delivered": False, "task_id": self.task_id, + "reason": "cancellation_waiting_for_child", "claim_released": True, + } if self.commit_delivery: self.history.append(marker( kwargs["sequence"], kwargs["call_kind"], sequence_span=kwargs["sequence_span"], @@ -152,6 +163,42 @@ async def test_claim_commits_delivery_and_reloads_before_cleanup( server.client.fail_workflow_task.assert_not_awaited() +async def test_pending_child_returns_to_other_work_and_replays_on_a_new_claim( + monkeypatch: pytest.MonkeyPatch, +) -> None: + @workflow.defn(name="other-work") + class OtherWorkflow: + def run(self, ctx: workflow.WorkflowContext, kind: str) -> str: + return kind + + server = ClaimServer() + server.pending_deliveries = 1 + worker = await server.worker(monkeypatch, workflows=[CancellationWorkflow, OtherWorkflow]) + assert await worker._run_workflow_task(claimed_task("child")) is None + assert server.trace == ["history", "delivery"] + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + other = claimed_task(observed=False) + other.update(task_id="other-task", workflow_id="other-workflow", run_id="other-run", workflow_type="other-work") + commands = await worker._run_workflow_task(other) + assert commands is not None and commands[0]["type"] == "complete_workflow" + + server.task_id = "parent-resume-task" + server.workflow_task_attempt = 1 + resumed = claimed_task("child") + resumed.update(task_id=server.task_id, workflow_task_attempt=1) + original_request = deepcopy(resumed["cancellation_request"]) + commands = await worker._run_workflow_task(resumed) + assert commands is not None and commands[0]["type"] == "start_timer" + assert commands[0]["delay_seconds"] == 1 + assert server.trace == ["history", "delivery", "completion", "history", "delivery", "history", "completion"] + assert resumed["cancellation_request"] == original_request + assert server.client.deliver_workflow_cancellation.await_args.kwargs["workflow_task_attempt"] == 1 + assert server.client.deliver_workflow_cancellation.await_args.kwargs["request_id"] == "request-1" + server.client.fail_workflow_task.assert_not_awaited() + + async def test_lost_delivery_ack_is_resolved_only_from_canonical_history(monkeypatch: pytest.MonkeyPatch) -> None: server = ClaimServer() server.delivery_error = TimeoutError("ack lost after commit") From 1763e72c02ce733d16b6fe735e353fa5e4e6472e Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 22:41:24 +0000 Subject: [PATCH 15/41] feat: expose immutable cancellation context in Python replay --- docs/cooperative-cancellation-design.md | 56 +++++ src/durable_workflow/__init__.py | 3 + .../_cooperative_cancellation.py | 39 +++- src/durable_workflow/cancellation.py | 138 ++++++++++++ src/durable_workflow/errors.py | 8 +- src/durable_workflow/workflow.py | 14 +- .../cooperative-cancellation-context.json | 17 ++ tests/test_cancellation_context.py | 203 ++++++++++++++++++ 8 files changed, 475 insertions(+), 3 deletions(-) create mode 100644 docs/cooperative-cancellation-design.md create mode 100644 src/durable_workflow/cancellation.py create mode 100644 tests/fixtures/cooperative-cancellation-context.json create mode 100644 tests/test_cancellation_context.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md new file mode 100644 index 0000000..d404d94 --- /dev/null +++ b/docs/cooperative-cancellation-design.md @@ -0,0 +1,56 @@ +# Cooperative cancellation source design + +This describes the unfinished source candidate for shared cancellation issue 136. +The default worker protocol remains 1.19. Cooperation requires explicit protocol +1.20 opt-in and a compatible Server and Native backend. The shared specification +and published mixed-language qualification remain pending. + +## Recorded context + +After committed cancellation delivery, `WorkflowContext.cancellation_context` +provides the original immutable request object. The delivered `WorkflowCancelled` +exception carries the same object in `context`. Earlier workflow code sees +`None`. A poll or heartbeat observation does not expose future cancellation +metadata to earlier workflow execution. + +```python +try: + yield ctx.start_timer(60) +except WorkflowCancelled as cancelled: + request = cancelled.context + with ctx.cancellation_shield(): + yield ctx.schedule_activity("release-reservation", [ + request.root_request_id if request else None, + request.reason if request else None, + ]) + raise +``` + +`CancellationContext` includes local and root request IDs, root workflow instance +and run IDs, the immediate parent request ID, original reason, requester, source, +root request time, original cleanup deadline and ordered lineage. Requester +metadata is limited to caller type, ID and label. It is a read-only mapping and +lineage is a tuple of immutable `CancellationLineage` entries. Timestamp values +are immutable UTC dates. `to_dict()` returns a detached portable snapshot. + +A child preserves the root request time and budget even when its local request +arrives later. Canonical request history supplies the context on cold replay, +while transport observation retains the opaque refresh route. The parser rejects +mismatched local request/run identities, cycles, invalid budgets and a delivery +that changes the accepted snapshot. Older cancellation histories without rich +context continue delivering cancellation with `context is None`. + +## Remaining qualification + +Rust context parity, portable operation policies, nested scopes and deterministic +remaining-time helpers still need completion. Remaining time must use the +replayed workflow clock. Do not subtract the host clock from the deadline in +workflow code. The runtime continues enforcing the original deadline and fencing +task and activity ownership. + +Connected qualification must cover the PHP parent, Python child, Rust remote +activity and PHP local activity together. Callbacks must stop without application +heartbeats. A replacement after SIGKILL during cleanup must replay the same +boundary and finish before the original 30-second deadline. Record supported +workflow lease, heartbeat and repair settings with that scenario. Exact published +artifacts and one cascade inspection view remain required for release claims. diff --git a/src/durable_workflow/__init__.py b/src/durable_workflow/__init__.py index adaed75..415a88d 100644 --- a/src/durable_workflow/__init__.py +++ b/src/durable_workflow/__init__.py @@ -16,6 +16,7 @@ AuthCompositionContractError, parse_auth_composition_contract, ) +from .cancellation import CancellationContext, CancellationLineage from .client import ( BridgeAdapterOutcome, Client, @@ -237,6 +238,8 @@ "ActivityInterceptorContext", "ActivityRetryPolicy", "BridgeAdapterOutcome", + "CancellationContext", + "CancellationLineage", "ChildWorkflowRetryPolicy", "ChildWorkflowCancelled", "ChildWorkflowFailed", diff --git a/src/durable_workflow/_cooperative_cancellation.py b/src/durable_workflow/_cooperative_cancellation.py index 34adbbf..332f5cc 100644 --- a/src/durable_workflow/_cooperative_cancellation.py +++ b/src/durable_workflow/_cooperative_cancellation.py @@ -11,6 +11,7 @@ from datetime import datetime from typing import Any +from .cancellation import CancellationContext from .errors import NonDeterministicReplayError REQUEST_EVENT = "CooperativeCancellationRequested" @@ -71,6 +72,7 @@ class CancellationRequest: requested_at: str cleanup_deadline_at: str history_refresh_page_token: str | None = None + context: CancellationContext | None = None @classmethod def from_observation(cls, value: Mapping[str, Any]) -> CancellationRequest: @@ -173,9 +175,33 @@ def read_cancellation_history( if kind == REQUEST_EVENT: if saw_request or delivery is not None: raise _invalid("history must contain one request before delivery") + context = None + if "cancellation" in payload: + snapshot = payload["cancellation"] + if not isinstance(snapshot, Mapping): + raise _invalid("canonical cancellation context must be an object") + try: + context = CancellationContext.from_dict(snapshot) + except ValueError as error: + raise _invalid(str(error)) from error + local = context.lineage[-1] + recorded_at = _timestamp(event.get("recorded_at", event.get("timestamp")), "recorded_at") + if ( + context.request_id != request_id or local.workflow_run_id != event_run + or ("workflow_instance_id" in payload + and local.workflow_instance_id != payload["workflow_instance_id"]) + or context.deadline != datetime.fromisoformat( + _timestamp(payload.get("cleanup_deadline_at"), "cleanup_deadline_at").replace("Z", "+00:00"), + ) + or context.requested_at > datetime.fromisoformat(recorded_at.replace("Z", "+00:00")) + or "reason" in payload and context.reason != payload["reason"] + ): + raise _invalid("canonical cancellation context does not match its request event") canonical_request = CancellationRequest( request_id, - request.requested_at + context.requested_at.isoformat(timespec="microseconds").replace("+00:00", "Z") + if context is not None + else request.requested_at if request is not None else _timestamp( event.get("recorded_at", event.get("timestamp")), @@ -183,6 +209,7 @@ def read_cancellation_history( ), _timestamp(payload.get("cleanup_deadline_at"), "cleanup_deadline_at"), request.history_refresh_page_token if request is not None else None, + context, ) if request is not None and ( request.request_id != canonical_request.request_id @@ -199,6 +226,16 @@ def read_cancellation_history( delivery = CancellationDelivery.from_payload(payload) if delivery.request_id != request.request_id: raise _invalid("delivery names a different request", delivery.sequence) + if "cancellation" in payload: + snapshot = payload["cancellation"] + if not isinstance(snapshot, Mapping) or request.context is None: + raise _invalid("delivery changes the canonical cancellation context", delivery.sequence) + try: + delivered_context = CancellationContext.from_dict(snapshot) + except ValueError as error: + raise _invalid(str(error), delivery.sequence) from error + if delivered_context != request.context: + raise _invalid("delivery changes the canonical cancellation context", delivery.sequence) delivery_index = index resolved: set[int] = set() failed: set[int] = set() diff --git a/src/durable_workflow/cancellation.py b/src/durable_workflow/cancellation.py new file mode 100644 index 0000000..7c12891 --- /dev/null +++ b/src/durable_workflow/cancellation.py @@ -0,0 +1,138 @@ +"""Immutable cancellation metadata recorded by the workflow runtime.""" + +from __future__ import annotations + +import re +from collections.abc import Mapping +from dataclasses import dataclass +from datetime import datetime, timezone +from types import MappingProxyType +from typing import Any + + +def _text(snapshot: Mapping[str, Any], key: str) -> str: + value = snapshot.get(key) + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"cancellation {key} must be a non-empty string") + return value + + +def _timestamp(value: str) -> datetime: + if not re.fullmatch(r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:\d{2})", value): + raise ValueError("cancellation timestamp requires an ISO date, time and timezone") + return datetime.fromisoformat(value.replace("Z", "+00:00")).astimezone(timezone.utc) + + +@dataclass(frozen=True) +class CancellationLineage: + """One local request in the ordered root-to-descendant lineage.""" + + request_id: str + workflow_instance_id: str + workflow_run_id: str + + def to_dict(self) -> dict[str, str]: + return { + "request_id": self.request_id, + "workflow_instance_id": self.workflow_instance_id, + "workflow_run_id": self.workflow_run_id, + } + + +@dataclass(frozen=True) +class CancellationContext: + """Original request metadata and budget, visible at committed delivery. + + Nested metadata is immutable. A child has a distinct local request ID + while retaining its root identity, request time and cleanup deadline. + """ + + request_id: str + root_request_id: str + root_workflow_instance_id: str + root_workflow_run_id: str + parent_request_id: str | None + reason: str | None + requester: Mapping[str, str] + source: str + requested_at: datetime + cleanup_deadline_at: datetime + lineage: tuple[CancellationLineage, ...] + + def __post_init__(self) -> None: + object.__setattr__(self, "requester", MappingProxyType(dict(self.requester))) + object.__setattr__(self, "lineage", tuple(self.lineage)) + + @property + def deadline(self) -> datetime: + """The original immutable cleanup deadline.""" + return self.cleanup_deadline_at + + @classmethod + def from_dict(cls, snapshot: Mapping[str, Any]) -> CancellationContext: + if snapshot.get("schema") != "durable-workflow.cancellation-context/v1": + raise ValueError("unsupported cancellation context schema") + requester = snapshot.get("requester") + if not isinstance(requester, Mapping) or not requester: + raise ValueError("cancellation requester must identify its caller") + normalized_requester: dict[str, str] = {} + for key, value in requester.items(): + if key not in {"type", "id", "label"} or not isinstance(value, str) or value == "": + raise ValueError("cancellation requester contains unsupported metadata") + normalized_requester[key] = value + lineage = snapshot.get("lineage") + if not isinstance(lineage, list) or not lineage: + raise ValueError("cancellation lineage must contain the root request") + normalized = [] + for entry in lineage: + if not isinstance(entry, Mapping): + raise ValueError("cancellation lineage entry is invalid") + normalized.append(CancellationLineage( + _text(entry, "request_id"), _text(entry, "workflow_instance_id"), _text(entry, "workflow_run_id"), + )) + request_id = _text(snapshot, "request_id") + root_request_id = _text(snapshot, "root_request_id") + root_instance_id = _text(snapshot, "root_workflow_instance_id") + root_run_id = _text(snapshot, "root_workflow_run_id") + parent_request_id = snapshot.get("parent_request_id") + reason = snapshot.get("reason") + if ( + parent_request_id is not None and (not isinstance(parent_request_id, str) or not parent_request_id) + or reason is not None and not isinstance(reason, str) + ): + raise ValueError("cancellation parent identity or reason is invalid") + if len({entry.request_id for entry in normalized}) != len(normalized) or len({ + entry.workflow_run_id for entry in normalized + }) != len(normalized): + raise ValueError("cancellation lineage cannot contain a cycle") + expected_parent = normalized[-2].request_id if len(normalized) > 1 else None + if ( + normalized[0] != CancellationLineage(root_request_id, root_instance_id, root_run_id) + or normalized[-1].request_id != request_id or parent_request_id != expected_parent + ): + raise ValueError("cancellation lineage does not match its request identities") + requested_at = _timestamp(_text(snapshot, "requested_at")) + deadline = _timestamp(_text(snapshot, "cleanup_deadline_at")) + if deadline <= requested_at: + raise ValueError("cancellation deadline must follow the original request") + return cls( + request_id, root_request_id, root_instance_id, root_run_id, parent_request_id, + reason, normalized_requester, _text(snapshot, "source"), requested_at, deadline, tuple(normalized), + ) + + def to_dict(self) -> dict[str, Any]: + """Return detached metadata in the portable context schema.""" + return { + "schema": "durable-workflow.cancellation-context/v1", + "request_id": self.request_id, + "root_request_id": self.root_request_id, + "root_workflow_instance_id": self.root_workflow_instance_id, + "root_workflow_run_id": self.root_workflow_run_id, + "parent_request_id": self.parent_request_id, + "reason": self.reason, + "requester": dict(self.requester), + "source": self.source, + "requested_at": self.requested_at.isoformat(timespec="microseconds").replace("+00:00", "Z"), + "cleanup_deadline_at": self.cleanup_deadline_at.isoformat(timespec="microseconds").replace("+00:00", "Z"), + "lineage": [entry.to_dict() for entry in self.lineage], + } diff --git a/src/durable_workflow/errors.py b/src/durable_workflow/errors.py index 143a2db..bb1f0cd 100644 --- a/src/durable_workflow/errors.py +++ b/src/durable_workflow/errors.py @@ -20,6 +20,8 @@ class explicitly (``except (ActivityCancelled, ...):``). This mirrors the way from typing import Any +from .cancellation import CancellationContext + class DurableWorkflowError(Exception): """Base class for every exception raised by the SDK.""" @@ -655,9 +657,13 @@ class WorkflowCancelled(BaseException): class by name. """ - def __init__(self, message: str = "workflow was cancelled", *, request_id: str | None = None) -> None: + def __init__( + self, message: str = "workflow was cancelled", *, request_id: str | None = None, + context: CancellationContext | None = None, + ) -> None: super().__init__(message) self.request_id = request_id + self.context = context class ActivityCancelled(BaseException): diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index aa6941d..cc835f8 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -34,6 +34,7 @@ from . import serializer from ._cooperative_cancellation import CancellationDelivery, read_cancellation_history +from .cancellation import CancellationContext from .client import WorkflowStreamAppendItem from .errors import ( ActivityFailed, @@ -1730,6 +1731,7 @@ def __init__( self._workflow_command_id = workflow_command_id or run_id or workflow_id self._cancel_requested = bool(cancel_requested) self._cancellation_request_id: str | None = None + self._cancellation_context: CancellationContext | None = None self._cancellation_shield_depth = 0 seed = int(hashlib.sha256(run_id.encode()).hexdigest()[:16], 16) self._rng = random.Random(seed) @@ -1824,7 +1826,15 @@ def is_cancellation_requested(self) -> bool: def throw_if_cancellation_requested(self) -> None: """Raise :class:`WorkflowCancelled` at an explicit safe point.""" if self._cancel_requested and self._cancellation_shield_depth == 0: - raise WorkflowCancelled("workflow cancellation was requested", request_id=self._cancellation_request_id) + raise WorkflowCancelled( + "workflow cancellation was requested", request_id=self._cancellation_request_id, + context=self._cancellation_context, + ) + + @property + def cancellation_context(self) -> CancellationContext | None: + """Original metadata, visible only at committed cancellation delivery.""" + return self._cancellation_context @contextlib.contextmanager def cancellation_shield(self) -> Generator[None, None, None]: @@ -5218,12 +5228,14 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl cancellation_consumed = True ctx._cancel_requested = True ctx._cancellation_request_id = boundary.request_id + ctx._cancellation_context = cancellation.request.context if cancellation.request is not None else None _apply_due_receivers() if isinstance(command, DurableOperationHandle): authored_sequence += 1 try: advanced_cmd = gen.throw(WorkflowCancelled( "workflow cancellation was requested", request_id=boundary.request_id, + context=ctx._cancellation_context, )) return None except StopIteration as stop: diff --git a/tests/fixtures/cooperative-cancellation-context.json b/tests/fixtures/cooperative-cancellation-context.json new file mode 100644 index 0000000..dbe7e89 --- /dev/null +++ b/tests/fixtures/cooperative-cancellation-context.json @@ -0,0 +1,17 @@ +{ + "schema": "durable-workflow.cancellation-context/v1", + "request_id": "request-1", + "root_request_id": "root-1", + "root_workflow_instance_id": "parent-instance", + "root_workflow_run_id": "parent-run", + "parent_request_id": "root-1", + "reason": "maintenance", + "requester": {"type": "operator", "id": "operator-1", "label": "Maintainer"}, + "source": "control_plane", + "requested_at": "2026-10-01T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-01T00:00:30.123456Z", + "lineage": [ + {"request_id": "root-1", "workflow_instance_id": "parent-instance", "workflow_run_id": "parent-run"}, + {"request_id": "request-1", "workflow_instance_id": "child-instance", "workflow_run_id": "run-1"} + ] +} diff --git a/tests/test_cancellation_context.py b/tests/test_cancellation_context.py new file mode 100644 index 0000000..2efe3cd --- /dev/null +++ b/tests/test_cancellation_context.py @@ -0,0 +1,203 @@ +from __future__ import annotations + +import json +from dataclasses import FrozenInstanceError +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import CancellationContext +from durable_workflow._cooperative_cancellation import read_cancellation_history +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.workflow import CompleteWorkflow, FailWorkflow, ScheduleActivity, WorkflowContext, replay +from tests.test_cooperative_cancellation import completed_activity + + +def snapshot() -> dict[str, Any]: + return json.loads((Path(__file__).parent / "fixtures/cooperative-cancellation-context.json").read_text()) + + +def request() -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationRequested", "recorded_at": "2026-10-01T00:00:05Z", + "payload": { + "workflow_command_id": "request-1", "workflow_instance_id": "child-instance", "workflow_run_id": "run-1", + "reason": "maintenance", "cleanup_deadline_at": "2026-10-01T00:00:30.123456Z", "cancellation": snapshot(), + }, + } + + +def delivery() -> dict[str, Any]: + return { + "event_type": "CooperativeCancellationDelivered", + "payload": { + "workflow_command_id": "request-1", "workflow_run_id": "run-1", "sequence": 1, + "call_kind": "timer", "cancellation": snapshot(), + }, + } + + +def observation() -> dict[str, Any]: + return { + "request_id": "request-1", "requested_at": "2026-10-01T00:00:05Z", + "cleanup_deadline_at": "2026-10-01T00:00:30.123456Z", "history_refresh_page_token": "opaque-first-page", + "cancellation": {**snapshot(), "reason": "untrusted observation"}, + } + + +def test_context_and_nested_metadata_are_immutable() -> None: + original = snapshot() + context = CancellationContext.from_dict(original) + original["reason"] = "changed" + original["requester"]["id"] = "changed" + original["lineage"][0]["request_id"] = "changed" + detached = context.to_dict() + detached["requester"]["id"] = "changed" + detached["lineage"].reverse() + assert context.to_dict() == snapshot() + assert context.request_id == "request-1" + assert context.root_request_id == context.parent_request_id == "root-1" + assert context.deadline == datetime(2026, 10, 1, 0, 0, 30, 123456, timezone.utc) + assert context.requested_at == datetime(2026, 10, 1, 0, 0, 0, 123456, timezone.utc) + with pytest.raises(FrozenInstanceError): + context.reason = "changed" # type: ignore[misc] + with pytest.raises(TypeError): + context.requester["id"] = "changed" # type: ignore[index] + with pytest.raises(FrozenInstanceError): + context.lineage[0].request_id = "changed" # type: ignore[misc] + + +def test_timezone_and_object_key_order_do_not_change_the_context() -> None: + value = snapshot() + value["requested_at"] = "2026-09-30T20:00:00.123456-04:00" + value["cleanup_deadline_at"] = "2026-09-30T20:00:30.123456-04:00" + value["requester"] = dict(reversed(list(value["requester"].items()))) + value["lineage"] = [dict(reversed(list(entry.items()))) for entry in value["lineage"]] + assert CancellationContext.from_dict(value) == CancellationContext.from_dict(snapshot()) + + +@pytest.mark.parametrize("field,value", [ + ("schema", "unknown"), ("request_id", ""), ("source", " "), ("parent_request_id", "wrong"), + ("root_request_id", "wrong"), ("root_workflow_run_id", "wrong"), ("reason", []), ("requester", []), + ("lineage", []), ("requested_at", "2026-02-30T00:00:00Z"), + ("cleanup_deadline_at", "2026-10-01T00:00:00.123456Z"), +]) +def test_invalid_context_is_rejected(field: str, value: Any) -> None: + with pytest.raises(ValueError): + CancellationContext.from_dict({**snapshot(), field: value}) + + +@pytest.mark.parametrize("field", ["request_id", "workflow_run_id"]) +def test_lineage_cannot_cycle(field: str) -> None: + value = snapshot() + value["lineage"][1][field] = value["lineage"][0][field] + with pytest.raises(ValueError, match="cycle"): + CancellationContext.from_dict(value) + + +def test_lineage_order_and_requester_metadata_are_validated() -> None: + value = snapshot() + value["lineage"].reverse() + with pytest.raises(ValueError, match="identities"): + CancellationContext.from_dict(value) + value = snapshot() + value["requester"]["authorization"] = "unsupported" + with pytest.raises(ValueError, match="unsupported metadata"): + CancellationContext.from_dict(value) + + +def test_canonical_child_keeps_root_time_after_expired_local_admission() -> None: + event = request() + event["recorded_at"] = "2026-10-01T00:00:31Z" + state = read_cancellation_history([event], run_id="run-1") + assert state.request is not None + assert state.request.requested_at == "2026-10-01T00:00:00.123456Z" + assert state.request.cleanup_deadline_at == "2026-10-01T00:00:30.123456Z" + assert state.request.context == CancellationContext.from_dict(snapshot()) + assert state.delivery is None + + +def test_observation_keeps_refresh_route_and_history_supplies_context() -> None: + observed = read_cancellation_history([], run_id="run-1", observation=observation()) + assert observed.request is not None and observed.request.context is None + state = read_cancellation_history([request(), delivery()], run_id="run-1", observation=observation()) + assert state.request is not None + assert state.request.context == CancellationContext.from_dict(snapshot()) + assert state.request.history_refresh_page_token == "opaque-first-page" + assert state.request.requested_at == "2026-10-01T00:00:00.123456Z" + + +@pytest.mark.parametrize("field,value", [ + ("workflow_command_id", "other"), ("workflow_instance_id", "other"), + ("cleanup_deadline_at", "2026-10-01T00:00:35Z"), ("reason", "changed"), ("cancellation", None), +]) +def test_context_must_match_canonical_local_request(field: str, value: Any) -> None: + event = request() + event["payload"][field] = value + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([event], run_id="run-1") + + +def test_context_must_name_local_run_and_not_postdate_admission() -> None: + event = request() + event["payload"]["cancellation"]["lineage"][1]["workflow_run_id"] = "other" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([event], run_id="run-1") + event = request() + event["recorded_at"] = "2026-09-30T23:59:59Z" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history([event], run_id="run-1") + + +def test_delivery_cannot_change_accepted_context() -> None: + event = delivery() + event["payload"]["cancellation"]["reason"] = "changed" + with pytest.raises(NonDeterministicReplayError, match="changes the canonical"): + read_cancellation_history([request(), event], run_id="run-1") + + +class ContextCleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + assert ctx.cancellation_context is None + try: + yield ctx.start_timer(10) + except WorkflowCancelled as cancelled: + assert cancelled.context is not None + assert cancelled.context is ctx.cancellation_context + with ctx.cancellation_shield(): + assert cancelled.context is ctx.cancellation_context + ctx.throw_if_cancellation_requested() + yield ctx.schedule_activity("cleanup", []) + return cancelled.context.to_dict() + return None + + +def test_cold_replay_exposes_same_context_only_at_delivery_and_cleanup() -> None: + history = [request(), delivery()] + first = replay(ContextCleanup, history, [], run_id="run-1") + assert len(first.commands) == 1 and isinstance(first.commands[0], ScheduleActivity) + assert first.commands[0].activity_type == "cleanup" + history.append(completed_activity(2, "cleanup", "cleaned")) + for _restart in range(2): + assert replay(ContextCleanup, history, [], run_id="run-1").commands == [CompleteWorkflow(snapshot())] + + +class ExplicitContextCheck: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(10) + except WorkflowCancelled: + try: + ctx.throw_if_cancellation_requested() + except WorkflowCancelled as cancelled: + assert cancelled.request_id == "request-1" + assert cancelled.context == CancellationContext.from_dict(snapshot()) + raise + + +def test_explicit_check_keeps_delivered_context() -> None: + outcome = replay(ExplicitContextCheck, [request(), delivery()], [], run_id="run-1") + assert len(outcome.commands) == 1 and isinstance(outcome.commands[0], FailWorkflow) + assert outcome.commands[0].exception_type == "WorkflowCancelled" From aeb54ff9e7ceb04dc2331dcba394a601e84d0854 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Thu, 1 Oct 2026 23:08:31 +0000 Subject: [PATCH 16/41] Preserve typed child cancellation policies in commands and cold replay --- docs/cooperative-cancellation-design.md | 32 ++- src/durable_workflow/__init__.py | 4 +- src/durable_workflow/cancellation.py | 38 +++ src/durable_workflow/worker.py | 21 ++ src/durable_workflow/workflow.py | 73 ++++- .../child-cancellation-policy-changed.json | 14 + tests/test_child_workflow_policies.py | 261 ++++++++++++++++++ tests/test_replay_regression_corpus.py | 7 + 8 files changed, 440 insertions(+), 10 deletions(-) create mode 100644 tests/fixtures/replay_regressions/child-cancellation-policy-changed.json create mode 100644 tests/test_child_workflow_policies.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index d404d94..d1e4627 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -40,9 +40,39 @@ mismatched local request/run identities, cycles, invalid budgets and a delivery that changes the accepted snapshot. Older cancellation histories without rich context continue delivering cancellation with `context is None`. +## Child policies + +`CancellationPolicy` and `ParentClosePolicy` are available from the package +root. Child commands accept these enums or their portable string values. + +```python +yield ctx.start_child_workflow( + "python.child", + [], + cancellation_policy=CancellationPolicy.WAIT_CANCELLATION_COMPLETED, + parent_close_policy=ParentClosePolicy.REQUEST_CANCELLATION, +) +``` + +`TRY_CANCEL` requests child cleanup and delivers parent cancellation without +waiting. `WAIT_CANCELLATION_COMPLETED` parks the parent until the child has a +recorded terminal outcome, releasing its task claim for other work. `ABANDON` +leaves the child independent and preserves the historical default. Parent +closure is a separate choice. `REQUEST_CANCELLATION` uses genuine cooperative +cleanup with the original lineage and budget. `REQUEST_CANCEL` retains legacy +terminal behavior. `TERMINATE` and `ABANDON` retain their existing meanings. + +Both command encoders preserve the policies. Cold replay compares them for +ordinary calls, parallel groups, selections and cancellation delivery. Omitted +historical fields mean the original `ABANDON` defaults. Later events with +missing fields keep the scheduled snapshot. Changed options and invalid or +conflicting history fail replay. A worker without the negotiated cooperation +capability refuses these choices before completion, reporting its identity and +the required protocol. Server also checks the immutable task claim and backend. + ## Remaining qualification -Rust context parity, portable operation policies, nested scopes and deterministic +Rust policy parity, portable activity policies, nested scopes and deterministic remaining-time helpers still need completion. Remaining time must use the replayed workflow clock. Do not subtract the host clock from the deadline in workflow code. The runtime continues enforcing the original deadline and fencing diff --git a/src/durable_workflow/__init__.py b/src/durable_workflow/__init__.py index 415a88d..398e3bb 100644 --- a/src/durable_workflow/__init__.py +++ b/src/durable_workflow/__init__.py @@ -16,7 +16,7 @@ AuthCompositionContractError, parse_auth_composition_contract, ) -from .cancellation import CancellationContext, CancellationLineage +from .cancellation import CancellationContext, CancellationLineage, CancellationPolicy, ParentClosePolicy from .client import ( BridgeAdapterOutcome, Client, @@ -240,6 +240,8 @@ "BridgeAdapterOutcome", "CancellationContext", "CancellationLineage", + "CancellationPolicy", + "ParentClosePolicy", "ChildWorkflowRetryPolicy", "ChildWorkflowCancelled", "ChildWorkflowFailed", diff --git a/src/durable_workflow/cancellation.py b/src/durable_workflow/cancellation.py index 7c12891..aeb0beb 100644 --- a/src/durable_workflow/cancellation.py +++ b/src/durable_workflow/cancellation.py @@ -6,10 +6,48 @@ from collections.abc import Mapping from dataclasses import dataclass from datetime import datetime, timezone +from enum import Enum from types import MappingProxyType from typing import Any +class CancellationPolicy(str, Enum): + """Cancellation at an awaiting child call. Abandon preserves the legacy default.""" + + TRY_CANCEL = "try_cancel" + WAIT_CANCELLATION_COMPLETED = "wait_cancellation_completed" + ABANDON = "abandon" + + +class ParentClosePolicy(str, Enum): + """Open-child behavior after parent closure. REQUEST_CANCEL is legacy terminal cancellation.""" + + ABANDON = "abandon" + REQUEST_CANCEL = "request_cancel" + REQUEST_CANCELLATION = "request_cancellation" + TERMINATE = "terminate" + + +def _canonical_child_policies(options: Mapping[str, Any]) -> dict[str, str]: + policies: dict[str, str] = {} + for field, enum in ( + ("parent_close_policy", ParentClosePolicy), + ("cancellation_policy", CancellationPolicy), + ): + value = options.get(field) + if value is None: + continue + if isinstance(value, Enum) and not isinstance(value, enum): + raise ValueError(f"child workflow {field} must be a supported policy") + if not isinstance(value, str): + raise ValueError(f"child workflow {field} must be a supported policy") + try: + policies[field] = enum(value).value + except ValueError as error: + raise ValueError(f"child workflow {field} must be a supported policy") from error + return policies + + def _text(snapshot: Mapping[str, Any], key: str) -> str: value = snapshot.get(key) if not isinstance(value, str) or not value.strip(): diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 4a74cd8..c2aa66f 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -93,6 +93,7 @@ RecordLocalActivity, RecordSideEffect, ReplayOutcome, + StartChildWorkflow, UpsertMemo, apply_update, commands_to_server_commands, @@ -2073,6 +2074,26 @@ def execute_local(command: RecordLocalActivity) -> Any: log.warning("failed to report workflow memo capability failure: %s", failure_error) return None + if not self._cooperative_cancellation_supported and any( + isinstance(command, StartChildWorkflow) and ( + command.parent_close_policy == "request_cancellation" + or command.cancellation_policy in ("try_cancel", "wait_cancellation_completed") + ) + for command in workflow_commands + ): + message = ( + f"child_cancellation_policy_not_supported: Python worker {self.worker_id} requires " + "cooperative_cancellation capability, worker protocol 1.20 and a compatible Server/Native backend" + ) + try: + await self.client.fail_workflow_task( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=attempt, + message=message, failure_type="RuntimeCapabilityUnsupported", + ) + except Exception as failure_error: + log.warning("failed to report child cancellation policy capability failure: %s", failure_error) + return None + def serialize_commands(source: list[Command]) -> list[dict[str, Any]]: return commands_to_server_commands( source, diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index cc835f8..a186e57 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -34,7 +34,7 @@ from . import serializer from ._cooperative_cancellation import CancellationDelivery, read_cancellation_history -from .cancellation import CancellationContext +from .cancellation import CancellationContext, CancellationPolicy, ParentClosePolicy, _canonical_child_policies from .client import WorkflowStreamAppendItem from .errors import ( ActivityFailed, @@ -918,10 +918,11 @@ class StartChildWorkflow: workflow_type: str arguments: list[Any] = field(default_factory=list) task_queue: str | None = None - parent_close_policy: str | None = None + parent_close_policy: str | ParentClosePolicy | None = None retry_policy: ChildWorkflowRetryPolicyInput | None = None execution_timeout_seconds: int | None = None run_timeout_seconds: int | None = None + cancellation_policy: str | CancellationPolicy | None = None _parallel_group_path: list[dict[str, Any]] | None = field( default=None, init=False, @@ -929,6 +930,12 @@ class StartChildWorkflow: compare=False, ) + def __post_init__(self) -> None: + _canonical_child_policies({ + "parent_close_policy": self.parent_close_policy, + "cancellation_policy": self.cancellation_policy, + }) + def to_server_command( self, task_queue: str, @@ -961,8 +968,10 @@ def to_server_command( cmd["queue"] = self.task_queue else: cmd["queue"] = task_queue - if self.parent_close_policy is not None: - cmd["parent_close_policy"] = self.parent_close_policy + cmd.update(_canonical_child_policies({ + "parent_close_policy": self.parent_close_policy, + "cancellation_policy": self.cancellation_policy, + })) if self.retry_policy is not None: cmd["retry_policy"] = ( self.retry_policy.to_dict() @@ -1491,8 +1500,10 @@ def commands_to_server_commands( task_queue=queue, ), )) - if command.parent_close_policy is not None: - server_command["parent_close_policy"] = command.parent_close_policy + server_command.update(_canonical_child_policies({ + "parent_close_policy": command.parent_close_policy, + "cancellation_policy": command.cancellation_policy, + })) if command.retry_policy is not None: server_command["retry_policy"] = ( command.retry_policy.to_dict() @@ -2080,7 +2091,8 @@ def start_child_workflow( arguments: list[Any] | None = None, *, task_queue: str | None = None, - parent_close_policy: str | None = None, + parent_close_policy: str | ParentClosePolicy | None = None, + cancellation_policy: str | CancellationPolicy | None = None, retry_policy: ChildWorkflowRetryPolicyInput | None = None, execution_timeout_seconds: int | None = None, run_timeout_seconds: int | None = None, @@ -2090,6 +2102,7 @@ def start_child_workflow( arguments=list(arguments) if arguments is not None else [], task_queue=task_queue, parent_close_policy=parent_close_policy, + cancellation_policy=cancellation_policy, retry_policy=retry_policy, execution_timeout_seconds=execution_timeout_seconds, run_timeout_seconds=run_timeout_seconds, @@ -3337,6 +3350,8 @@ def _recorded_step_details(payload: Mapping[str, Any]) -> dict[str, Any]: for key in ( "workflow_type", "child_workflow_type", + "parent_close_policy", + "cancellation_policy", "timer_kind", "change_id", "condition_key", @@ -3486,6 +3501,20 @@ def _recorded_detail_mismatch(command: Any, step: _RecordedStep) -> str | None: f"Recorded child workflow_type {recorded!r}, but current workflow " f"started {command.workflow_type!r}." ) + actual_policies = { + "parent_close_policy": "abandon", "cancellation_policy": "abandon", + **_canonical_child_policies({ + "parent_close_policy": command.parent_close_policy, + "cancellation_policy": command.cancellation_policy, + }), + } + for field, actual_policy in actual_policies.items(): + recorded_policy = step.details.get(field, "abandon") + if recorded_policy != actual_policy: + return ( + f"child_workflow_policy_changed: recorded {field} {recorded_policy!r}, " + f"but current workflow requested {actual_policy!r}." + ) elif isinstance(command, RecordVersionMarker): recorded = step.details.get("change_id") if isinstance(recorded, str) and recorded != command.change_id: @@ -3662,6 +3691,7 @@ def _replay_state( workflow_id = workflow_id or _workflow_id_from_history(events) event_types_by_sequence: dict[int, list[str]] = {} details_by_sequence: dict[int, dict[str, Any]] = {} + child_policies_by_sequence: dict[int, dict[str, str]] = {} resolved_sequences: set[int] = set() condition_wait_ids_by_sequence: dict[int, str] = {} selected_condition_wait_ids_by_sequence: dict[int, str] = {} @@ -3683,6 +3713,29 @@ def _replay_state( # envelopes into the public non-determinism diagnostic. recorded_details = {} details_by_sequence.setdefault(sequence, {}).update(recorded_details) + if event_type in ( + "ChildWorkflowScheduled", "ChildRunStarted", "ChildRunCompleted", + "ChildRunFailed", "ChildRunCancelled", "ChildRunTerminated", + ): + try: + incoming_policies = _canonical_child_policies(payload) + except ValueError as error: + raise NonDeterministicReplayError( + sequence, "child workflow", [event_type], + detail=f"invalid_child_workflow_policy_history: {error}", + ) from error + previous_policies = child_policies_by_sequence.get(sequence) + for field, value in incoming_policies.items(): + if previous_policies is not None and previous_policies[field] != value: + raise NonDeterministicReplayError( + sequence, "child workflow", [event_type], + detail=f"child_workflow_policy_history_conflict: {field} changed between history events.", + ) + child_policies_by_sequence[sequence] = { + "parent_close_policy": "abandon", "cancellation_policy": "abandon", + **(previous_policies or {}), **incoming_policies, + } + details_by_sequence[sequence].update(child_policies_by_sequence[sequence]) if event_type == "ConditionWaitOpened": wait_id = payload.get("condition_wait_id") if isinstance(wait_id, str) and wait_id: @@ -3870,6 +3923,8 @@ def _recorded_step( else {} ) details.update(_recorded_step_details(payload)) + if shape == "child workflow": + details.update(child_policies_by_sequence.get(workflow_sequence, {})) return _RecordedStep( workflow_sequence=workflow_sequence, shape=shape, @@ -5196,7 +5251,9 @@ def _assert_cancellation_call_matches(command: Any, boundary: CancellationDelive continue _assert_step_matches(leaf, _RecordedStep( workflow_sequence=sequence, shape=opening_shapes[event_type], - event_types=[event_type], details=_recorded_step_details(payload), + event_types=[event_type], details={ + **details_by_sequence.get(sequence, {}), **_recorded_step_details(payload), + }, )) def _terminal_state(value: Any, *, include_pending: bool) -> _ReplayState: diff --git a/tests/fixtures/replay_regressions/child-cancellation-policy-changed.json b/tests/fixtures/replay_regressions/child-cancellation-policy-changed.json new file mode 100644 index 0000000..9de5b32 --- /dev/null +++ b/tests/fixtures/replay_regressions/child-cancellation-policy-changed.json @@ -0,0 +1,14 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "python-child-cancellation-policy-changed", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": {"type": "tests.replay.child-policy-author", "input": [], "payload_codec": "avro"}, + "history": [ + {"event_type": "ChildWorkflowScheduled", "payload": {"sequence": 1, "child_workflow_type": "child", "parent_close_policy": "abandon", "cancellation_policy": "wait_cancellation_completed"}}, + {"event_type": "ChildRunCompleted", "payload": {"sequence": 1, "child_workflow_type": "child"}} + ], + "expected": {"command_sequence": []}, + "expected_replay_error": {"type": "NonDeterministicReplayError", "message_contains": "child_workflow_policy_changed", "workflow_sequence": 1} +} diff --git a/tests/test_child_workflow_policies.py b/tests/test_child_workflow_policies.py new file mode 100644 index 0000000..8fb1250 --- /dev/null +++ b/tests/test_child_workflow_policies.py @@ -0,0 +1,261 @@ +from __future__ import annotations + +from typing import Any + +import pytest + +from durable_workflow import CancellationPolicy, ParentClosePolicy, serializer +from durable_workflow import workflow as workflow_module +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import ( + CompleteWorkflow, + StartChildWorkflow, + WorkflowContext, + commands_to_server_commands, + replay, +) +from tests.test_cooperative_cancellation import marker, request +from tests.test_cooperative_cancellation_worker import ClaimServer, claimed_task + + +def options() -> dict[str, Any]: + return { + "parent_close_policy": ParentClosePolicy.REQUEST_CANCELLATION, + "cancellation_policy": CancellationPolicy.WAIT_CANCELLATION_COMPLETED, + } + + +def workflow(mode: str, policies: dict[str, Any]) -> type: + class PolicyWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + first = ctx.start_child_workflow("child", ["argument"], **policies) + try: + if mode == "parallel": + return (yield [first]) + if mode == "selection": + return (yield ctx.select({ + "first": first, + "second": ctx.start_child_workflow("other-child", [], **policies), + })) + return (yield first) + except WorkflowCancelled as cancelled: + return cancelled.request_id + + return PolicyWorkflow + + +def scheduled(commands: list[Any]) -> list[dict[str, Any]]: + return [ + {"event_type": "ChildWorkflowScheduled", "payload": { + "sequence": index + 1, "child_workflow_type": command["workflow_type"], + "child_workflow_run_id": f"child-run-{index + 1}", **command, + }} + for index, command in enumerate(commands_to_server_commands(commands, "queue")) + ] + + +@pytest.mark.parametrize("parent", list(ParentClosePolicy)) +@pytest.mark.parametrize("operation", list(CancellationPolicy)) +def test_typed_and_string_policies_round_trip_through_both_encoders_and_cold_replay( + parent: ParentClosePolicy, operation: CancellationPolicy, +) -> None: + child = StartChildWorkflow("child", ["argument"], parent_close_policy=parent, cancellation_policy=operation) + direct = child.to_server_command("queue") + assert direct == commands_to_server_commands([child], "queue")[0] + assert direct == StartChildWorkflow( + "child", ["argument"], parent_close_policy=parent.value, cancellation_policy=operation.value, + ).to_server_command("queue") + assert type(direct["parent_close_policy"]) is str + assert type(direct["cancellation_policy"]) is str + assert serializer.decode_envelope(direct["arguments"]) == ["argument"] + history = scheduled([child]) + [{"event_type": "ChildRunCompleted", "payload": { + "sequence": 1, "child_workflow_type": "child", "result": serializer.envelope("recorded-result"), + }}] + result = replay(workflow("sequential", { + "parent_close_policy": parent, "cancellation_policy": operation, + }), history, []) + assert result.commands == [CompleteWorkflow("recorded-result")] + + +@pytest.mark.parametrize("field", ["parent_close_policy", "cancellation_policy"]) +@pytest.mark.parametrize("value", ["unknown", "", True, 1, [], {"type": "abandon"}]) +def test_invalid_options_are_rejected_before_scheduling(field: str, value: Any) -> None: + with pytest.raises(ValueError, match=field): + StartChildWorkflow("child", **{field: value}) + + +def test_policy_enum_types_are_not_interchangeable() -> None: + with pytest.raises(ValueError, match="parent_close_policy"): + StartChildWorkflow("child", parent_close_policy=CancellationPolicy.ABANDON) + with pytest.raises(ValueError, match="cancellation_policy"): + StartChildWorkflow("child", cancellation_policy=ParentClosePolicy.ABANDON) + + +def test_existing_positional_constructor_and_omitted_wire_fields_are_preserved() -> None: + child = StartChildWorkflow("child", ["argument"], "queue", "request_cancel", None, 60, 30) + assert child.execution_timeout_seconds == 60 + assert child.run_timeout_seconds == 30 + assert child.cancellation_policy is None + assert "cancellation_policy" not in child.to_server_command("queue") + omitted = StartChildWorkflow("child").to_server_command("queue") + assert "parent_close_policy" not in omitted + assert "cancellation_policy" not in omitted + result = replay(workflow("sequential", { + "parent_close_policy": ParentClosePolicy.ABANDON, "cancellation_policy": CancellationPolicy.ABANDON, + }), scheduled([StartChildWorkflow("child", ["argument"])]), []) + assert len(result.commands) == 1 + assert isinstance(result.commands[0], StartChildWorkflow) + assert result.commands[0].cancellation_policy == CancellationPolicy.ABANDON + + +@pytest.mark.parametrize("mode", ["sequential", "parallel", "selection"]) +@pytest.mark.parametrize("field", ["parent_close_policy", "cancellation_policy"]) +@pytest.mark.parametrize("delivered", [False, True]) +def test_changed_policies_fail_cold_replay_and_cancellation_delivery(mode: str, field: str, delivered: bool) -> None: + policies = options() + history = scheduled(replay(workflow(mode, policies), [], []).commands) + if delivered: + history += [request(), marker(1, "child" if mode == "sequential" else "parallel", sequence_span=len(history))] + assert replay(workflow(mode, policies), history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + else: + cold = replay(workflow(mode, policies), history, []) + if mode == "selection": + assert cold.commands == [] + else: + assert scheduled(cold.commands) == history + policies[field] = "abandon" + with pytest.raises(NonDeterministicReplayError, match="child_workflow_policy_changed") as captured: + replay(workflow(mode, policies), history, [], run_id="run-1") + assert captured.value.workflow_sequence == 1 + + +@pytest.mark.parametrize("mode", ["sequential", "parallel", "selection"]) +@pytest.mark.parametrize("field", ["parent_close_policy", "cancellation_policy"]) +def test_historical_defaults_cannot_change_to_cooperative_policies(mode: str, field: str) -> None: + history = scheduled(replay(workflow(mode, {}), [], []).commands) + with pytest.raises(NonDeterministicReplayError, match="child_workflow_policy_changed"): + replay(workflow(mode, {field: options()[field]}), history, []) + + +@pytest.mark.parametrize("event", [ + "ChildRunStarted", "ChildRunCompleted", "ChildRunFailed", "ChildRunCancelled", "ChildRunTerminated", +]) +@pytest.mark.parametrize("policy,reason", [ + ("try_cancel", "child_workflow_policy_history_conflict"), + ({"type": "try_cancel"}, "invalid_child_workflow_policy_history"), +]) +def test_invalid_and_conflicting_child_history_is_rejected(event: str, policy: Any, reason: str) -> None: + history = scheduled(replay(workflow("sequential", options()), [], []).commands) + history.append({"event_type": event, "payload": {"sequence": 1, "cancellation_policy": policy}}) + with pytest.raises(NonDeterministicReplayError, match=reason): + replay(workflow("sequential", options()), history, []) + + +def test_later_start_and_terminal_events_preserve_the_scheduled_snapshot() -> None: + history = scheduled(replay(workflow("sequential", options()), [], []).commands) + history += [ + {"event_type": "ChildRunStarted", "payload": {"sequence": 1, "child_workflow_type": "child"}}, + {"event_type": "ChildRunCompleted", "payload": {"sequence": 1, "result": serializer.envelope("done")}}, + ] + assert replay(workflow("sequential", options()), history, []).commands == [CompleteWorkflow("done")] + + +def test_selection_handle_keeps_the_loser_policy_after_the_winner_and_during_cancellation() -> None: + def authoring(second_policy: CancellationPolicy) -> type: + class AwaitSecond: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + selected = yield ctx.select({ + "first": ctx.start_child_workflow("child", ["argument"], **options()), + "second": ctx.start_child_workflow("other-child", [], **{ + **options(), "cancellation_policy": second_policy, + }), + }) + try: + return (yield selected.handles["second"].await_result()) + except WorkflowCancelled as cancelled: + return cancelled.request_id + + return AwaitSecond + + original = authoring(CancellationPolicy.WAIT_CANCELLATION_COMPLETED) + history = scheduled(replay(original, [], [], run_id="run-1").commands) + history += [ + {"id": "child-completed", "event_type": "ChildRunCompleted", "payload": { + **history[0]["payload"], "result": serializer.envelope("first-result"), + }}, + {"event_type": "SelectionResolved", "payload": { + "selection_group_id": "select-calls:1:2", "selection_group_base_sequence": 1, "selection_group_size": 2, + "member_key": "first", "member_index": 0, "member_base_sequence": 1, "member_size": 1, + "operation_kind": "child", "operation_identity": "child-run-1", "outcome": "completed", + "resolution_event_id": "child-completed", "resolution_event_type": "ChildRunCompleted", + }}, + ] + assert replay(original, history, [], run_id="run-1").commands == [] + for delivered in (False, True): + current = list(history) + if delivered: + current += [request(), marker(3, "selection_handle", operation_sequence=2, operation_sequence_span=1)] + assert replay(original, current, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + with pytest.raises(NonDeterministicReplayError, match="child_workflow_policy_changed") as captured: + replay(authoring(CancellationPolicy.TRY_CANCEL), current, [], run_id="run-1") + assert captured.value.workflow_sequence == 2 + + +@pytest.mark.parametrize("protocol", ["1.19", "1.20"]) +@pytest.mark.parametrize("policies", [ + {"parent_close_policy": ParentClosePolicy.REQUEST_CANCELLATION}, + {"cancellation_policy": CancellationPolicy.TRY_CANCEL}, + {"cancellation_policy": CancellationPolicy.WAIT_CANCELLATION_COMPLETED}, +]) +async def test_worker_without_opt_in_refuses_cooperative_child_commands( + monkeypatch: pytest.MonkeyPatch, protocol: str, policies: dict[str, Any], +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", protocol) + cls = workflow_module.defn(name="child-policy-worker")(workflow("sequential", policies)) + server = ClaimServer(history=[]) + worker = Worker(server.client, task_queue="queue", worker_id="cooperative-worker", workflows=[cls]) + await worker._register() + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "child-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is None + server.client.complete_workflow_task.assert_not_awaited() + failure = server.client.fail_workflow_task.await_args.kwargs + assert failure["failure_type"] == "RuntimeCapabilityUnsupported" + assert "child_cancellation_policy_not_supported" in failure["message"] + assert "cooperative-worker" in failure["message"] + assert "1.20" in failure["message"] + + +@pytest.mark.parametrize("parent", [ + ParentClosePolicy.ABANDON, ParentClosePolicy.REQUEST_CANCEL, ParentClosePolicy.TERMINATE, +]) +async def test_legacy_child_policies_remain_available_without_cooperation( + monkeypatch: pytest.MonkeyPatch, parent: ParentClosePolicy, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.19") + cls = workflow_module.defn(name="child-policy-worker")(workflow("sequential", { + "parent_close_policy": parent, "cancellation_policy": CancellationPolicy.ABANDON, + })) + server = ClaimServer(history=[]) + worker = Worker(server.client, task_queue="queue", worker_id="cooperative-worker", workflows=[cls]) + await worker._register() + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "child-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is not None and result[0]["parent_close_policy"] == parent.value + server.client.fail_workflow_task.assert_not_awaited() + + +async def test_capable_worker_transmits_both_typed_child_policies(monkeypatch: pytest.MonkeyPatch) -> None: + cls = workflow_module.defn(name="child-policy-worker")(workflow("sequential", options())) + server = ClaimServer(history=[]) + worker = await server.worker(monkeypatch, workflows=[cls]) + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "child-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is not None + assert result[0]["parent_close_policy"] == "request_cancellation" + assert result[0]["cancellation_policy"] == "wait_cancellation_completed" + server.client.fail_workflow_task.assert_not_awaited() diff --git a/tests/test_replay_regression_corpus.py b/tests/test_replay_regression_corpus.py index 9875162..972cebd 100644 --- a/tests/test_replay_regression_corpus.py +++ b/tests/test_replay_regression_corpus.py @@ -224,7 +224,14 @@ def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] return "not cancelled" +@workflow.defn(name="tests.replay.child-policy-author") +class ChildPolicyAuthorWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + return (yield ctx.start_child_workflow("child", [])) + + WORKFLOWS = [ + ChildPolicyAuthorWorkflow, ColdReplacementSatisfiedConditionWorkflow, CooperativeReopenedConditionCleanupWorkflow, GoldenSagaCompensationWorkflow, From c17ee56af2ade4b94f0a5e137bf3779acbc251f1 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 01:22:34 +0000 Subject: [PATCH 17/41] Validate original remote cancellation stop receipts in the Python client --- docs/cooperative-cancellation-design.md | 15 ++ src/durable_workflow/client.py | 51 ++++++ tests/test_cooperative_cancellation_client.py | 151 ++++++++++++++++++ 3 files changed, 217 insertions(+) diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index d1e4627..0f00052 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -70,6 +70,21 @@ conflicting history fail replay. A worker without the negotiated cooperation capability refuses these choices before completion, reporting its identity and the required protocol. Server also checks the immutable task claim and backend. +## Remote callback-stop transport + +`Client.acknowledge_activity_cancellation()` reports the original task, activity +attempt, lease owner and cancellation request. It requires explicit worker +protocol 1.20 and validates the Server's original receipt identity. Duplicate +retries retain those identities and share a five-second transport budget. +Refusals preserve the Server diagnostic. This receipt does not renew authority, +record an application heartbeat or extend the cleanup deadline. + +The caller must first prove the callback stopped and was joined. The current +Python worker fences abandoned callback threads but cannot forcibly stop them. +It therefore does not send this receipt. Independent process supervision is +the next implementation step. Cooperating downstream systems still need +idempotency or reconciliation for effects already performed. + ## Remaining qualification Rust policy parity, portable activity policies, nested scopes and deterministic diff --git a/src/durable_workflow/client.py b/src/durable_workflow/client.py index 311a087..c22cc0c 100644 --- a/src/durable_workflow/client.py +++ b/src/durable_workflow/client.py @@ -5590,6 +5590,57 @@ async def activity_task_status( json={"activity_attempt_id": activity_attempt_id, "lease_owner": lease_owner}, timeout=5.0, ), timeout=5.0) + async def acknowledge_activity_cancellation( + self, + *, + task_id: str, + activity_attempt_id: str, + lease_owner: str, + request_id: str, + ) -> dict[str, Any]: + """Report that the original owner's remote callback has stopped. + + The caller must stop and join the callback before calling this method. + The receipt records diagnostic evidence and does not renew a lease, + heartbeat, publication authority or cleanup budget. Requires explicit + worker protocol 1.20. Retries share one five-second transport budget. + """ + if not _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )): + raise ValueError("activity cancellation acknowledgment requires explicit worker protocol 1.20") + identities = { + "task_id": task_id, + "activity_attempt_id": activity_attempt_id, + "lease_owner": lease_owner, + "request_id": request_id, + } + if any( + not isinstance(value, str) or not value.strip() or len(value.encode("utf-8")) > 255 + for value in identities.values() + ): + raise ValueError( + "activity cancellation acknowledgment requires bounded, nonempty claim and request identities" + ) + result = await asyncio.wait_for(self._request( + "POST", f"/worker/activity-tasks/{quote(task_id, safe='._:-')}/acknowledge-cancellation", + worker=True, + json={"activity_attempt_id": activity_attempt_id, "lease_owner": lease_owner, "request_id": request_id}, + timeout=5.0, + ), timeout=5.0) + if ( + not isinstance(result, dict) + or any(result.get(key) != value for key, value in identities.items()) + or result.get("acknowledged") is not True + or not isinstance(result.get("duplicate"), bool) + or result.get("reason") is not None + or result.get("heartbeat_recorded") is not False + or not isinstance(result.get("history_event_id"), str) + or not result["history_event_id"].strip() + ): + raise ServerError(200, {"reason": "invalid_activity_cancellation_acknowledgement"}) + return result + async def heartbeat_activity_task( self, *, diff --git a/tests/test_cooperative_cancellation_client.py b/tests/test_cooperative_cancellation_client.py index 476002e..8925799 100644 --- a/tests/test_cooperative_cancellation_client.py +++ b/tests/test_cooperative_cancellation_client.py @@ -1,5 +1,7 @@ from __future__ import annotations +import asyncio +import time from collections.abc import AsyncIterator from typing import Any from unittest.mock import AsyncMock, patch @@ -45,6 +47,14 @@ def pending_delivery_response(**overrides: Any) -> dict[str, Any]: } +def activity_stop_response(**overrides: Any) -> dict[str, Any]: + return { + "task_id": "task/1", "activity_attempt_id": "attempt-1", "lease_owner": "worker-1", + "request_id": "original-request", "acknowledged": True, "duplicate": False, + "reason": None, "heartbeat_recorded": False, "history_event_id": "original-event", **overrides, + } + + @pytest_asyncio.fixture async def client() -> AsyncIterator[Client]: async with Client( @@ -322,6 +332,147 @@ async def test_active_claim_refusal_does_not_fall_back_to_immediate_cancel(clien assert send.call_args.args[1].endswith("/request-cancellation") +@pytest.mark.asyncio +async def test_activity_stop_receipt_uses_original_identity_and_worker_credential( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[response(activity_stop_response()), + response(activity_stop_response(duplicate=True))]) as send: + original = await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + duplicate = await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert original["history_event_id"] == duplicate["history_event_id"] == "original-event" + assert duplicate["duplicate"] is True + for call in send.await_args_list: + assert call.args[:2] == ("POST", "/api/worker/activity-tasks/task%2F1/acknowledge-cancellation") + assert call.kwargs["headers"]["Authorization"] == "Bearer worker-token" + assert call.kwargs["headers"]["X-Namespace"] == "ns1" + assert call.kwargs["headers"]["X-Durable-Workflow-Protocol-Version"] == "1.20" + assert call.kwargs["json"] == { + "activity_attempt_id": "attempt-1", "lease_owner": "worker-1", "request_id": "original-request", + } + assert call.kwargs["timeout"] == 5.0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [ + {"task_id": "other"}, {"activity_attempt_id": "other"}, {"lease_owner": "other"}, {"request_id": "other"}, + {"acknowledged": 1}, {"acknowledged": False}, {"duplicate": 1}, {"reason": "refused"}, + {"heartbeat_recorded": True}, {"heartbeat_recorded": 0}, {"history_event_id": None}, {"history_event_id": " "}, +]) +async def test_activity_stop_receipt_rejects_mismatched_or_unproved_response( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict[str, Any], +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response(activity_stop_response(**change))), + pytest.raises(ServerError) as error, + ): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert error.value.reason() == "invalid_activity_cancellation_acknowledgement" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ["task_id", "activity_attempt_id", "lease_owner", "request_id"]) +@pytest.mark.parametrize("identity", [None, True, "", " ", "x" * 256, "é" * 128]) +async def test_activity_stop_receipt_rejects_invalid_identity_before_io( + client: Client, monkeypatch: pytest.MonkeyPatch, field: str, identity: Any, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + options = {"task_id": "task/1", "activity_attempt_id": "attempt-1", "lease_owner": "worker-1", + "request_id": "original-request", field: identity} + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.acknowledge_activity_cancellation(**options) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("version", [None, "1.19", "2.20", "invalid"]) +async def test_activity_stop_receipt_requires_explicit_worker_opt_in( + client: Client, monkeypatch: pytest.MonkeyPatch, version: str | None, +) -> None: + if version is None: + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", raising=False) + else: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", version) + with patch.object(client._http, "request", new_callable=AsyncMock) as send, pytest.raises(ValueError): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + send.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("status", "reason"), [(409, "lease_owner_mismatch"), (503, "storage_fenced")]) +async def test_activity_stop_receipt_preserves_server_refusal( + client: Client, monkeypatch: pytest.MonkeyPatch, status: int, reason: str, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with ( + patch.object(client._http, "request", new_callable=AsyncMock, + return_value=response({"reason": reason, "request_admitted": False}, status)) as send, + pytest.raises(ServerError) as error, + ): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert error.value.reason() == reason + assert send.await_count == (1 if status == 409 else 3) + assert all(call == send.await_args_list[0] for call in send.await_args_list) + + +@pytest.mark.asyncio +async def test_activity_stop_receipt_retries_original_identity_after_response_loss( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + with patch.object(client._http, "request", new_callable=AsyncMock, + side_effect=[httpx.ReadTimeout("receipt response lost"), + response(activity_stop_response(duplicate=True))]) as send: + result = await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert result["duplicate"] is True + assert result["history_event_id"] == "original-event" + assert send.await_count == 2 + assert send.await_args_list[0] == send.await_args_list[1] + + +@pytest.mark.asyncio +async def test_activity_stop_receipt_transport_has_one_total_budget( + client: Client, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + stopped = asyncio.Event() + + async def blocked(*args: Any, **kwargs: Any) -> httpx.Response: + try: + await asyncio.Event().wait() + finally: + stopped.set() + raise AssertionError("blocked transport resumed") + + started = time.monotonic() + with ( + patch.object(client._http, "request", new_callable=AsyncMock, side_effect=blocked) as send, + pytest.raises(asyncio.TimeoutError), + ): + await client.acknowledge_activity_cancellation( + task_id="task/1", activity_attempt_id="attempt-1", lease_owner="worker-1", request_id="original-request", + ) + assert time.monotonic() - started < 6.0 + assert stopped.is_set() + assert send.await_count == 1 + + @pytest.mark.asyncio @pytest.mark.parametrize("reason", ["lease_owner_mismatch", "workflow_task_attempt_mismatch", "lease_expired"]) async def test_delivery_refusal_preserves_the_server_lease_reason( From 07b2fb0f8348613cf8741e5e796cca81439a8a6b Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 01:28:52 +0000 Subject: [PATCH 18/41] Own cooperative activity callbacks in independent spawn supervisors --- docs/cooperative-cancellation-design.md | 22 +- src/durable_workflow/_activity_process.py | 306 ++++++++++++++++++++++ tests/test_activity_process.py | 229 ++++++++++++++++ 3 files changed, 554 insertions(+), 3 deletions(-) create mode 100644 src/durable_workflow/_activity_process.py create mode 100644 tests/test_activity_process.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index 0f00052..8b2b1e5 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -81,9 +81,25 @@ record an application heartbeat or extend the cleanup deadline. The caller must first prove the callback stopped and was joined. The current Python worker fences abandoned callback threads but cannot forcibly stop them. -It therefore does not send this receipt. Independent process supervision is -the next implementation step. Cooperating downstream systems still need -idempotency or reconciliation for effects already performed. +It therefore does not send this receipt. The internal process supervisor is +implemented for the next worker integration step. Each attempt gets an explicit +spawn context, an independent supervisor, a callback process and private payload +files. Single-byte control channels keep owner disconnect independent of a +payload transfer or callback progress. The supervisor joins the callback before +reporting stop. The owner then joins the supervisor. A failed supervisor alone +does not prove a live callback stopped. + +The primitive's process tests cover a C call holding the callback interpreter's +GIL, ignored TERM followed by forced stop, actual owner SIGKILL, typed results, +application failure metadata, interceptors and authored heartbeats. It is not +connected to worker claims or Server receipts yet. Those tests remain required +before worker adoption. Handlers, arguments and interceptors must be compatible +with Python's spawn serialization. Define importable handlers and protect the +application entry point with `if __name__ == "__main__"`. Captured memory changes +are local to the callback process. Open process-local connections in the +callback. Legacy worker protocol 1.19 continues using its existing execution. +Cooperating downstream systems still need idempotency or reconciliation for +effects already performed. ## Remaining qualification diff --git a/src/durable_workflow/_activity_process.py b/src/durable_workflow/_activity_process.py new file mode 100644 index 0000000..0891961 --- /dev/null +++ b/src/durable_workflow/_activity_process.py @@ -0,0 +1,306 @@ +"""Owned callback processes for the unfinished cooperative worker. + +The supervisor never unpickles or invokes the serialized application callback. +Single-byte control messages keep owner disconnect observable even if an +application payload is large or its writer is killed. Attempt-local private +files carry payloads. This module does not grant or report Server authority. +""" + +from __future__ import annotations + +import asyncio +import inspect +import multiprocessing +import pickle +import selectors +import shutil +import socket +import tempfile +import threading +import traceback +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from multiprocessing.process import BaseProcess +from pathlib import Path +from typing import Any, cast + +from .activity import ActivityContext, ActivityInfo, _set_context +from .errors import ActivityCancelled, NonRetryableError +from .interceptors import ActivityInterceptorContext, WorkerInterceptor + + +class CallbackProcessLost(RuntimeError): + """The owner could not prove the callback outcome or physical stop.""" + + +@dataclass(frozen=True) +class CallbackInvocation: + handler: Callable[..., Any] + args: tuple[Any, ...] + info: ActivityInfo + task: dict[str, Any] + interceptors: tuple[WorkerInterceptor, ...] = () + + def serialize(self) -> bytes: + try: + return pickle.dumps(self, protocol=pickle.HIGHEST_PROTOCOL) + except Exception as error: + raise ValueError( + f"cooperative activity {self.info.activity_type!r} requires a spawn-compatible " + "importable handler, arguments and interceptors" + ) from error + + +@dataclass(frozen=True) +class CallbackFailure: + message: str + failure_type: str + failure_class: str + failure_code: int | None + stack_trace: str + non_retryable: bool + + +@dataclass(frozen=True) +class CallbackOutcome: + value: Any = None + failure: CallbackFailure | None = None + + +def _write_payload(directory: str, name: str, payload: Any) -> None: + destination = Path(directory, name) + temporary = destination.with_suffix(".writing") + with temporary.open("wb") as stream: + pickle.dump(payload, stream, protocol=pickle.HIGHEST_PROTOCOL) + temporary.replace(destination) + + +def _read_payload(directory: str, name: str) -> Any: + with Path(directory, name).open("rb") as stream: + return pickle.load(stream) + + +def _callback(directory: str, channel: socket.socket) -> None: + # Deserialize only in the process that is permitted to run application code. + heartbeat_lock = threading.Lock() + + def heartbeat_request(details: dict[str, Any] | None) -> None: + with heartbeat_lock: + _write_payload(directory, "heartbeat", details) + channel.sendall(b"H") + if channel.recv(1) != b"Y": + raise ActivityCancelled() + + async def heartbeat(details: dict[str, Any] | None) -> None: + await asyncio.to_thread(heartbeat_request, details) + + async def invoke(invocation: CallbackInvocation) -> Any: + _set_context(ActivityContext(info=invocation.info, client=cast(Any, None), heartbeat_callback=heartbeat)) + context = ActivityInterceptorContext( + worker_id=invocation.info.worker_id, task_queue=invocation.info.task_queue, + task=invocation.task, activity_type=invocation.info.activity_type, args=invocation.args, + ) + + async def call(ctx: ActivityInterceptorContext) -> Any: + value = (invocation.handler(*ctx.args) if inspect.iscoroutinefunction(invocation.handler) + else await asyncio.to_thread(invocation.handler, *ctx.args)) + return await value if inspect.isawaitable(value) else value + + handler = call + for interceptor in reversed(invocation.interceptors): + next_handler = handler + + async def intercepted( + ctx: ActivityInterceptorContext, *, interceptor: WorkerInterceptor = interceptor, + next_handler: Callable[[ActivityInterceptorContext], Awaitable[Any]] = next_handler, + ) -> Any: + return await interceptor.execute_activity(ctx, next_handler) + + handler = intercepted + try: + return await handler(context) + finally: + _set_context(None) + + try: + invocation = pickle.loads(Path(directory, "invocation").read_bytes()) + if not isinstance(invocation, CallbackInvocation): + raise TypeError("invalid callback invocation") + channel.sendall(b"B") + try: + _write_payload(directory, "outcome", CallbackOutcome(value=asyncio.run(invoke(invocation)))) + except BaseException as error: + code = getattr(error, "code", None) + _write_payload(directory, "outcome", CallbackOutcome(failure=CallbackFailure( + message=str(error), failure_type=type(error).__name__, + failure_class=f"{type(error).__module__}.{type(error).__qualname__}", + failure_code=code if isinstance(code, int) and not isinstance(code, bool) else None, + stack_trace=traceback.format_exc(), non_retryable=isinstance(error, NonRetryableError), + ))) + finally: + channel.close() + + +def _stop_owned_callback(callback: BaseProcess) -> None: + if callback.is_alive(): + callback.terminate() + callback.join(timeout=0.5) + if callback.is_alive(): + callback.kill() + # If the OS cannot reap it yet, retain the supervisor and ownership. Never + # send stop evidence before an actual join, even after the owner disappears. + callback.join() + + +def _supervise(directory: str, owner: socket.socket) -> None: + context = multiprocessing.get_context("spawn") + supervisor_channel, callback_channel = socket.socketpair() + callback = context.Process(target=_callback, args=(directory, callback_channel)) + callback_started = False + stopped = False + try: + callback.start() + callback_started = True + callback_channel.close() + with selectors.DefaultSelector() as selector: + selector.register(owner, selectors.EVENT_READ, "owner") + selector.register(supervisor_channel, selectors.EVENT_READ, "callback") + while True: + for key, _ in selector.select(timeout=0.05): + message = cast(socket.socket, key.fileobj).recv(1) + if key.data == "owner": + if message in (b"", b"F"): + return + if message == b"S": + _stop_owned_callback(callback) + stopped = True + owner.sendall(b"X") + elif message in (b"Y", b"N") and not stopped: + supervisor_channel.sendall(message) + elif message in (b"B", b"H"): + owner.sendall(message) + elif message == b"": + selector.unregister(supervisor_channel) + if not stopped and callback.exitcode is not None: + callback.join() + stopped = True + owner.sendall(b"D") + except (BrokenPipeError, ConnectionResetError): + pass # Owner disappeared. The finally block still owns and reaps its callback. + finally: + if callback_started: + _stop_owned_callback(callback) + callback.close() + callback_channel.close() + supervisor_channel.close() + owner.close() + shutil.rmtree(directory, ignore_errors=True) + + +class SupervisedCallback: + """One owner, one supervisor and one isolated application callback. + + The owner must serialize access to result()/stop(). Cancellation of result() + leaves the process owned until stop() proves shutdown. A killed supervisor + never supplies stop evidence. The worker integration remains a separate gate. + """ + + def __init__(self, invocation: CallbackInvocation) -> None: + serialized = invocation.serialize() + self.directory = tempfile.mkdtemp(prefix="dw-activity-") + try: + Path(self.directory, "invocation").write_bytes(serialized) + except BaseException: + shutil.rmtree(self.directory, ignore_errors=True) + raise + self._owner, self._supervisor_channel = socket.socketpair() + self._owner.setblocking(False) + self._supervisor = multiprocessing.get_context("spawn").Process( + target=_supervise, args=(self.directory, self._supervisor_channel), + ) + self._started = False + self._closed = False + self._joined = False + self._exitcode: int | None = None + self.stopped = False + + async def start(self) -> None: + try: + # Pass only the private directory and control socket to spawn. The + # payload is already serialized, and this short start cannot leave + # a background spawn thread untracked if the owner is cancelled. + self._supervisor.start() + self._started = True + self._supervisor_channel.close() + message = await asyncio.wait_for(self._receive(), timeout=5.0) + if message != b"B": + raise CallbackProcessLost("callback did not prove spawn readiness") + except BaseException: + await self.close() + raise + + async def _receive(self) -> bytes: + try: + message = await asyncio.get_running_loop().sock_recv(self._owner, 1) + except (ConnectionResetError, OSError) as error: + raise CallbackProcessLost("callback supervisor disconnected") from error + if not message: + raise CallbackProcessLost("callback supervisor disconnected without stop evidence") + return message + + async def _send(self, message: bytes) -> None: + await asyncio.get_running_loop().sock_sendall(self._owner, message) + + async def result(self, heartbeat: Callable[[dict[str, Any] | None], Awaitable[None]]) -> CallbackOutcome: + while True: + message = await self._receive() + if message == b"H": + await heartbeat(_read_payload(self.directory, "heartbeat")) + await self._send(b"Y") + elif message == b"D": + outcome = _read_payload(self.directory, "outcome") + if not isinstance(outcome, CallbackOutcome): + raise CallbackProcessLost("callback did not produce a typed outcome") + await self._finish() + return outcome + else: + raise CallbackProcessLost("callback supervisor returned an unexpected result boundary") + + async def stop(self) -> None: + if self.stopped: + return + async def wait_for_stop() -> None: + while await self._receive() not in (b"D", b"X"): + pass + await self._finish() + + try: + await self._send(b"S") + await asyncio.wait_for(wait_for_stop(), timeout=5.0) + except BaseException: + await self.close() + raise + + async def _finish(self) -> None: + await self._send(b"F") + await self.close() + if self._exitcode != 0: + raise CallbackProcessLost("callback supervisor did not confirm a clean join") + self.stopped = True + + async def close(self) -> None: + if self._joined: + return + if not self._closed: + self._closed = True + self._owner.close() + self._supervisor_channel.close() + if self._started: + await asyncio.to_thread(self._supervisor.join, 5.0) + if self._supervisor.is_alive(): + raise CallbackProcessLost("callback supervisor has not confirmed shutdown") + self._exitcode = self._supervisor.exitcode + else: + shutil.rmtree(self.directory, ignore_errors=True) + self._supervisor.close() + self._joined = True diff --git a/tests/test_activity_process.py b/tests/test_activity_process.py new file mode 100644 index 0000000..82e3d14 --- /dev/null +++ b/tests/test_activity_process.py @@ -0,0 +1,229 @@ +from __future__ import annotations + +import asyncio +import ctypes +import os +import signal +import subprocess +import sys +import time +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import activity +from durable_workflow._activity_process import ( + CallbackInvocation, + CallbackProcessLost, + SupervisedCallback, +) +from durable_workflow.activity import ActivityInfo +from durable_workflow.errors import NonRetryableError +from durable_workflow.interceptors import ActivityHandler, ActivityInterceptorContext, PassthroughWorkerInterceptor + + +def invocation(handler: Any, *args: Any, interceptors: tuple[Any, ...] = ()) -> CallbackInvocation: + return CallbackInvocation( + handler=handler, args=args, + info=ActivityInfo("task", "process-test", "attempt", 1, "queue", "owner"), + task={"task_id": "task", "activity_attempt_id": "attempt"}, interceptors=interceptors, + ) + + +def typed_result(value: bytes) -> dict[str, Any]: + return {"bytes": value, "task": activity.context().info.task_id, "nested": [1, None, {"x": True}]} + + +async def authored_heartbeat() -> bytes: + await activity.context().heartbeat({"progress": b"typed"}) + return b"result" + + +def authored_sync_heartbeat() -> bytes: + return asyncio.run(authored_heartbeat()) + + +class ProcessInterceptor(PassthroughWorkerInterceptor): + async def execute_activity(self, context: ActivityInterceptorContext, next: ActivityHandler) -> Any: + assert context.worker_id == "owner" + assert context.task["activity_attempt_id"] == "attempt" + return {"intercepted": await next(context)} + + +class ProcessFailure(NonRetryableError): + code = 42 + + +def failing() -> None: + raise ProcessFailure("original failure") + + +async def blocked_without_python_progress(marker: str) -> None: + signal.signal(signal.SIGTERM, signal.SIG_IGN) + Path(marker).write_text(str(os.getpid())) + # PyDLL retains the callback interpreter's GIL during the C call. A Python + # thread or signal callback in this interpreter cannot supervise this work. + ctypes.PyDLL(None).sleep(60) + + +def wait_for_release(marker: str) -> None: + Path(marker).write_text(str(os.getpid())) + while not Path(marker + ".release").exists(): + time.sleep(0.02) + + +async def wait_for_file(path: Path) -> None: + async def ready() -> None: + while not path.exists(): + await asyncio.sleep(0.02) + await asyncio.wait_for(ready(), timeout=5.0) + + +async def wait_for_exit(pid: int) -> None: + async def gone() -> None: + while True: + try: + os.kill(pid, 0) + except ProcessLookupError: + return + await asyncio.sleep(0.02) + await asyncio.wait_for(gone(), timeout=7.0) + + +async def no_heartbeat(details: dict[str, Any] | None) -> None: + raise AssertionError("callback emitted an unauthored heartbeat") + + +async def test_spawn_preserves_typed_result_context_and_interceptors() -> None: + callback = SupervisedCallback(invocation(typed_result, b"\x00\xff", interceptors=(ProcessInterceptor(),))) + try: + await callback.start() + outcome = await callback.result(no_heartbeat) + assert outcome.failure is None + assert outcome.value == {"intercepted": {"bytes": b"\x00\xff", "task": "task", + "nested": [1, None, {"x": True}]}} + assert callback.stopped is True + assert not Path(callback.directory).exists() + finally: + await callback.close() + + +@pytest.mark.parametrize("handler", [authored_heartbeat, authored_sync_heartbeat]) +async def test_only_authored_heartbeat_crosses_to_the_owner(handler: Any) -> None: + observed: list[Any] = [] + owner_loop = asyncio.get_running_loop() + + async def heartbeat(details: dict[str, Any] | None) -> None: + assert asyncio.get_running_loop() is owner_loop + observed.append(details) + + callback = SupervisedCallback(invocation(handler)) + try: + await callback.start() + outcome = await callback.result(heartbeat) + assert outcome.value == b"result" + assert outcome.failure is None + assert observed == [{"progress": b"typed"}] + assert callback.stopped is True + finally: + await callback.close() + + +async def test_spawn_preserves_application_failure_metadata() -> None: + callback = SupervisedCallback(invocation(failing)) + try: + await callback.start() + outcome = await callback.result(no_heartbeat) + assert outcome.failure is not None + assert outcome.failure.message == "original failure" + assert outcome.failure.failure_type == "ProcessFailure" + assert outcome.failure.failure_class == "tests.test_activity_process.ProcessFailure" + assert outcome.failure.failure_code == 42 + assert outcome.failure.non_retryable is True + assert "raise ProcessFailure" in outcome.failure.stack_trace + assert callback.stopped is True + finally: + await callback.close() + + +@pytest.mark.skipif(os.name != "posix", reason="POSIX process signals") +async def test_stop_kills_and_joins_a_gil_blocked_callback_without_heartbeats(tmp_path: Path) -> None: + marker = tmp_path / "callback" + callback = SupervisedCallback(invocation(blocked_without_python_progress, str(marker))) + try: + await callback.start() + await wait_for_file(marker) + pid = int(marker.read_text()) + started = time.monotonic() + await callback.stop() + assert time.monotonic() - started < 3.0 + assert callback.stopped is True + await wait_for_exit(pid) + assert not Path(callback.directory).exists() + finally: + await callback.close() + + +@pytest.mark.skipif(os.name != "posix", reason="POSIX process signals") +async def test_owner_sigkill_leaves_supervisor_to_stop_and_reap_callback(tmp_path: Path) -> None: + marker = tmp_path / "callback" + supervisor_marker = tmp_path / "supervisor" + script = tmp_path / "owner.py" + script.write_text("""import asyncio +import sys +from pathlib import Path +from durable_workflow._activity_process import SupervisedCallback +from tests.test_activity_process import blocked_without_python_progress, invocation +async def run(): + callback = SupervisedCallback(invocation(blocked_without_python_progress, sys.argv[1])) + await callback.start() + Path(sys.argv[2]).write_text(str(callback._supervisor.pid) + '\\n' + callback.directory) + await asyncio.Event().wait() +if __name__ == '__main__': + asyncio.run(run()) +""") + environment = {**os.environ, "PYTHONPATH": os.pathsep.join(sys.path)} + owner = subprocess.Popen([sys.executable, str(script), str(marker), str(supervisor_marker)], env=environment) + try: + await wait_for_file(marker) + await wait_for_file(supervisor_marker) + supervisor_pid, directory = supervisor_marker.read_text().splitlines() + callback_pid = int(marker.read_text()) + os.kill(owner.pid, signal.SIGKILL) + await asyncio.to_thread(owner.wait, 5.0) + assert owner.returncode == -signal.SIGKILL + await wait_for_exit(callback_pid) + await wait_for_exit(int(supervisor_pid)) + assert not Path(directory).exists() + finally: + if owner.poll() is None: + owner.kill() + await asyncio.to_thread(owner.wait, 5.0) + + +@pytest.mark.skipif(os.name != "posix", reason="POSIX process signals") +async def test_dead_supervisor_with_live_callback_never_proves_stop(tmp_path: Path) -> None: + marker = tmp_path / "callback" + callback = SupervisedCallback(invocation(wait_for_release, str(marker))) + try: + await callback.start() + await wait_for_file(marker) + pid = int(marker.read_text()) + callback._supervisor.kill() + with pytest.raises((CallbackProcessLost, BrokenPipeError, ConnectionResetError)): + await callback.stop() + assert callback.stopped is False + os.kill(pid, 0) # A dead supervisor alone left the callback alive. + finally: + Path(str(marker) + ".release").touch() + if marker.exists(): + await wait_for_exit(int(marker.read_text())) + await callback.close() + + +def test_nonimportable_handler_is_rejected_before_spawn() -> None: + def callback() -> None: + pass + with pytest.raises(ValueError, match="spawn-compatible.*importable handler"): + SupervisedCallback(invocation(callback)) From 16ae1f0bed5cdad58c4ce1ddde4f4ce8c6157b0c Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 01:34:26 +0000 Subject: [PATCH 19/41] Retain exact Native overlay provenance in Python source qualification --- .github/workflows/ci.yml | 55 +++++++++++++++++++++--- README.md | 6 ++- docker-compose.native-cancellation.yml | 14 ++++++ scripts/ci/checkout-public-repository.py | 1 + tests/test_ci_checkout.py | 12 ++++-- 5 files changed, 78 insertions(+), 10 deletions(-) create mode 100644 docker-compose.native-cancellation.yml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index cea83cb..1182035 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -15,6 +15,10 @@ on: description: Qualify the candidate cooperative protocol 1.20 type: boolean default: false + native_commit: + description: Optional exact public Native source overlay SHA for candidate qualification + type: string + default: "" permissions: contents: read @@ -157,6 +161,8 @@ jobs: env: DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION: ${{ inputs.cooperative_qualification && '1.20' || '1.19' }} DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION: ${{ inputs.cooperative_qualification && '1' || '0' }} + DURABLE_WORKFLOW_NATIVE_SOURCE: ${{ github.workspace }}/native + COMPOSE_FILE: docker-compose.test.yml steps: - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: @@ -168,6 +174,29 @@ jobs: env: SERVER_COMMIT: ${{ inputs.server_commit }} run: python sdk-python/scripts/ci/checkout-public-repository.py server server --commit "$SERVER_COMMIT" + - name: Require exact cooperative candidate identities + env: + COOPERATIVE_REQUESTED: ${{ inputs.cooperative_qualification }} + SERVER_COMMIT: ${{ inputs.server_commit }} + NATIVE_COMMIT: ${{ inputs.native_commit }} + run: | + if [ "$COOPERATIVE_REQUESTED" = true ]; then + [[ "$SERVER_COMMIT" =~ ^[0-9a-f]{40}$ ]] + fi + if [ -n "$NATIVE_COMMIT" ]; then + test "$COOPERATIVE_REQUESTED" = true + [[ "$NATIVE_COMMIT" =~ ^[0-9a-f]{40}$ ]] + fi + - name: Check out optional exact public Native source + if: ${{ inputs.native_commit != '' }} + env: + NATIVE_COMMIT: ${{ inputs.native_commit }} + run: python sdk-python/scripts/ci/checkout-public-repository.py workflow native --commit "$NATIVE_COMMIT" + - name: Select readonly Native source qualification + if: ${{ inputs.native_commit != '' }} + run: | + printf 'COMPOSE_FILE=docker-compose.test.yml:docker-compose.native-cancellation.yml\n' >> "$GITHUB_ENV" + printf 'DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION=1\n' >> "$GITHUB_ENV" - run: pip install -e '.[dev]' working-directory: sdk-python - name: Configure isolated Docker project @@ -176,22 +205,38 @@ jobs: - name: Start Server stack working-directory: sdk-python run: | - docker compose --project-name "$COMPOSE_PROJECT_NAME" -f docker-compose.test.yml \ + docker compose --project-name "$COMPOSE_PROJECT_NAME" \ up -d --build --wait --timeout 300 + - name: Retain exact source and published image authority + working-directory: sdk-python + env: + NATIVE_COMMIT: ${{ inputs.native_commit }} + run: | + docker compose exec -T server cat /app/.package-provenance > integration-package-provenance.txt + docker compose images --format json > integration-images.json + jq --null-input --arg sdk "$GITHUB_SHA" --arg server "$(git -C ../server rev-parse HEAD)" --arg native "$NATIVE_COMMIT" \ + '{qualification:"source",sdk_commit:$sdk,server_commit:$server,native_source_overlay:($native | if length > 0 then . else null end)}' \ + > integration-source-provenance.json - name: Select and probe integration endpoint working-directory: sdk-python run: python scripts/ci/configure-integration-endpoint.py - name: Run integration tests + shell: bash working-directory: sdk-python env: DURABLE_WORKFLOW_AUTH_TOKEN: test-token - run: pytest tests/integration/ -v --junitxml=integration-results.xml + run: pytest tests/integration/ -v --capture=tee-sys --junitxml=integration-results.xml 2>&1 | tee integration-scenarios.log - name: Retain connected scenario results if: ${{ always() && github.server_url == 'https://github.com' }} uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7 with: name: connected-integration-${{ github.sha }} - path: sdk-python/integration-results.xml + path: | + sdk-python/integration-results.xml + sdk-python/integration-scenarios.log + sdk-python/integration-source-provenance.json + sdk-python/integration-package-provenance.txt + sdk-python/integration-images.json if-no-files-found: warn retention-days: 7 - name: Emit integration diagnostics @@ -200,7 +245,7 @@ jobs: run: | for service in bootstrap server worker repair; do echo "::group::$service logs" - docker compose --project-name "$COMPOSE_PROJECT_NAME" -f docker-compose.test.yml \ + docker compose --project-name "$COMPOSE_PROJECT_NAME" \ logs --no-color --tail 200 "$service" || true echo "::endgroup::" done @@ -208,7 +253,7 @@ jobs: if: always() working-directory: sdk-python run: | - docker compose --project-name "$COMPOSE_PROJECT_NAME" -f docker-compose.test.yml \ + docker compose --project-name "$COMPOSE_PROJECT_NAME" \ down -v --rmi local target-branch-qualification: diff --git a/README.md b/README.md index 1641a15..76a4b3f 100644 --- a/README.md +++ b/README.md @@ -137,7 +137,11 @@ docker compose -f docker-compose.test.yml down -v Candidate cooperative cancellation qualification is explicit. In a manual CI run, supply an exact public `server_commit` and set `cooperative_qualification` to true. CI verifies that checkout, builds the candidate Server, enables protocol -1.20, runs the connected cases and retains their JUnit results. For a local +1.20, runs the connected cases and retains JUnit, raw observations, image +authority and exact source provenance. An optional exact `native_commit` mounts +that public Native checkout read-only into the test stack. The image's published +Composer authority stays intact and the evidence identifies the source overlay. +These are source qualification runs. For a local candidate, set `DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION=1.20` before starting Compose and `DURABLE_WORKFLOW_COOPERATIVE_QUALIFICATION=1` for pytest. These cases fail if the runtime does not discover the required capability. Ordinary CI diff --git a/docker-compose.native-cancellation.yml b/docker-compose.native-cancellation.yml new file mode 100644 index 0000000..6fffe46 --- /dev/null +++ b/docker-compose.native-cancellation.yml @@ -0,0 +1,14 @@ +# Source qualification only. Published Composer authority stays in the image. +services: + bootstrap: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro + server: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro + worker: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro + repair: + volumes: + - ${DURABLE_WORKFLOW_NATIVE_SOURCE:?exact Native checkout}:/app/vendor/durable-workflow/workflow:ro diff --git a/scripts/ci/checkout-public-repository.py b/scripts/ci/checkout-public-repository.py index 12758bf..3b15ff8 100644 --- a/scripts/ci/checkout-public-repository.py +++ b/scripts/ci/checkout-public-repository.py @@ -13,6 +13,7 @@ PUBLIC_REPOSITORIES = { "cli": "https://github.com/durable-workflow/cli.git", "server": "https://github.com/durable-workflow/server.git", + "workflow": "https://github.com/durable-workflow/workflow.git", } diff --git a/tests/test_ci_checkout.py b/tests/test_ci_checkout.py index 637bd27..66ec059 100644 --- a/tests/test_ci_checkout.py +++ b/tests/test_ci_checkout.py @@ -22,6 +22,7 @@ [ ("cli", "https://github.com/durable-workflow/cli.git"), ("server", "https://github.com/durable-workflow/server.git"), + ("workflow", "https://github.com/durable-workflow/workflow.git"), ], ) def test_public_checkout_uses_github_authority_on_every_runner( @@ -90,7 +91,10 @@ def test_candidate_checkout_rejects_non_sha_before_running_git(tmp_path: Path, c @pytest.mark.parametrize("matches", [True, False]) -def test_candidate_checkout_verifies_the_requested_public_commit(tmp_path: Path, matches: bool) -> None: +@pytest.mark.parametrize("repository", ["server", "workflow"]) +def test_candidate_checkout_verifies_the_requested_public_commit( + tmp_path: Path, matches: bool, repository: str, +) -> None: commit = "a" * 40 resolved = commit if matches else "b" * 40 capture = tmp_path / "git-calls" @@ -104,13 +108,13 @@ def test_candidate_checkout_verifies_the_requested_public_commit(tmp_path: Path, **os.environ, "PATH": str(tmp_path), "GIT_CAPTURE": str(capture), "RESOLVED_COMMIT": resolved, } result = subprocess.run( - [sys.executable, str(CHECKOUT_SCRIPT), "server", str(tmp_path / "server"), "--commit", commit], + [sys.executable, str(CHECKOUT_SCRIPT), repository, str(tmp_path / repository), "--commit", commit], env=environment, capture_output=True, text=True, ) assert (result.returncode == 0) is matches calls = capture.read_text().splitlines() - assert "https://github.com/durable-workflow/server.git" in calls[0] + assert f"https://github.com/durable-workflow/{repository}.git" in calls[0] assert any(f"fetch --depth=1 origin {commit}" in call for call in calls) assert all("credential.helper=" in call for call in calls) if matches: - assert f"Integration server source commit: {commit}" in result.stdout + assert f"Integration {repository} source commit: {commit}" in result.stdout From 19febbf704a64ca557e9877e6b04dd443690b86b Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 01:50:11 +0000 Subject: [PATCH 20/41] Stop and join cooperative Python remote callbacks before acknowledging cancellation --- docs/cooperative-cancellation-design.md | 30 +- src/durable_workflow/_activity_process.py | 2 + src/durable_workflow/worker.py | 199 +++++++----- tests/test_cooperative_cancellation_worker.py | 4 + tests/test_cooperative_remote_worker.py | 293 +++++++++++------- 5 files changed, 342 insertions(+), 186 deletions(-) diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index 8b2b1e5..bbd1fa9 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -79,25 +79,37 @@ retries retain those identities and share a five-second transport budget. Refusals preserve the Server diagnostic. This receipt does not renew authority, record an application heartbeat or extend the cleanup deadline. -The caller must first prove the callback stopped and was joined. The current -Python worker fences abandoned callback threads but cannot forcibly stop them. -It therefore does not send this receipt. The internal process supervisor is -implemented for the next worker integration step. Each attempt gets an explicit +The cooperative remote worker stops and joins its callback and supervisor +before sending this receipt. Each attempt gets an explicit spawn context, an independent supervisor, a callback process and private payload files. Single-byte control channels keep owner disconnect independent of a payload transfer or callback progress. The supervisor joins the callback before reporting stop. The owner then joins the supervisor. A failed supervisor alone does not prove a live callback stopped. -The primitive's process tests cover a C call holding the callback interpreter's +The process tests cover a C call holding the callback interpreter's GIL, ignored TERM followed by forced stop, actual owner SIGKILL, typed results, -application failure metadata, interceptors and authored heartbeats. It is not -connected to worker claims or Server receipts yet. Those tests remain required -before worker adoption. Handlers, arguments and interceptors must be compatible +application failure metadata, interceptors and authored heartbeats. The worker +observes ownership independently of callback progress, checks before result +encoding or failure reporting, and reports only the original canonical request +after confirmed stop. Unconfirmed stop retains activity capacity and refuses +successful worker shutdown. A receipt refusal cannot become result publication +or a new cleanup budget. Connected exact-source qualification remains required. + +After application drain expires, cooperative remote shutdown permits up to +15 seconds for process reaping and bounded receipt transport. Application work +is stopped during this phase. The run's original cancellation deadline stays +unchanged. Failure to confirm stop leaves the worker registration active and +raises an explicit shutdown error. + +Handlers, arguments and interceptors must be compatible with Python's spawn serialization. Define importable handlers and protect the application entry point with `if __name__ == "__main__"`. Captured memory changes are local to the callback process. Open process-local connections in the -callback. Legacy worker protocol 1.19 continues using its existing execution. +callback. Registration and remote polling refuse incompatible handler or +interceptor definitions before claiming work, naming the worker and activity. +Cooperative local callbacks still require their separate supervision and receipt +model. Legacy worker protocol 1.19 continues using its existing execution. Cooperating downstream systems still need idempotency or reconciliation for effects already performed. diff --git a/src/durable_workflow/_activity_process.py b/src/durable_workflow/_activity_process.py index 0891961..39f4289 100644 --- a/src/durable_workflow/_activity_process.py +++ b/src/durable_workflow/_activity_process.py @@ -59,6 +59,7 @@ class CallbackFailure: failure_code: int | None stack_trace: str non_retryable: bool + cancelled: bool = False @dataclass(frozen=True) @@ -136,6 +137,7 @@ async def intercepted( failure_class=f"{type(error).__module__}.{type(error).__qualname__}", failure_code=code if isinstance(code, int) and not isinstance(code, bool) else None, stack_trace=traceback.format_exc(), non_retryable=isinstance(error, NonRetryableError), + cancelled=isinstance(error, ActivityCancelled), ))) finally: channel.close() diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index c2aa66f..86738b6 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -22,6 +22,7 @@ import inspect import json import logging +import pickle import sys import threading import time @@ -36,6 +37,7 @@ from typing import Annotated, Any, Concatenate, Literal, ParamSpec, TypeVar, Union, get_args, get_origin, get_type_hints from . import serializer +from ._activity_process import CallbackFailure, CallbackInvocation, CallbackProcessLost, SupervisedCallback from ._cooperative_cancellation import CancellationRequest, read_cancellation_history from .activity import ActivityContext, ActivityInfo, _set_context from .auth_composition import ( @@ -164,6 +166,12 @@ class _RemoteActivityExecutionAborted(Exception): """Ownership cannot authorize another remote callback boundary.""" +class _RemoteActivityApplicationFailure(Exception): + def __init__(self, failure: CallbackFailure) -> None: + super().__init__(failure.message) + self.failure = failure + + class _LocalActivityTimedOut(Exception): def __init__(self, kind: str) -> None: super().__init__(f"local activity {kind} timeout elapsed") @@ -1071,6 +1079,7 @@ def __init__( self._local_activity_executor: ThreadPoolExecutor | None = None self._remote_activity_executor: ThreadPoolExecutor | None = None self._remote_activity_threads: set[Future[Any]] = set() + self._remote_activity_processes: set[SupervisedCallback] = set() self._remote_activity_thread_slots = asyncio.Semaphore(max_concurrent_activity_tasks) self._wf_semaphore = asyncio.Semaphore(max_concurrent_workflow_tasks) self._act_semaphore = asyncio.Semaphore(max_concurrent_activity_tasks) @@ -1216,6 +1225,7 @@ async def _register(self) -> None: ) if "cooperative_cancellation" in self.capabilities and not self._cooperative_cancellation_supported: raise RuntimeError("cooperative cancellation requires explicit compatible runtime and worker protocol 1.20") + self._validate_cooperative_activity_handlers() self._query_tasks_supported = _server_supports_query_tasks(info) self._workflow_memo_updates_supported = _server_supports_workflow_memo_updates(info) has_update_validators = any( @@ -2214,6 +2224,53 @@ async def _report_workflow_task_after_completion_error( fail_error, ) + def _validate_cooperative_activity_handlers(self) -> None: + if not self._cooperative_cancellation_supported: + return + for activity_type, handler in self.activities.items(): + try: + pickle.dumps((handler, self.interceptors), protocol=pickle.HIGHEST_PROTOCOL) + except Exception as error: + raise RuntimeError( + f"worker {self.worker_id!r} cannot supervise cooperative activity {activity_type!r}: " + "handler and interceptors must be importable and spawn-compatible" + ) from error + + async def _acknowledge_stopped_remote_activity(self, task: dict[str, Any]) -> None: + # The caller has joined both the callback and its supervisor. Observation + # binds only this original claim to canonical cancellation metadata. + task_id = task["task_id"] + attempt_id = task.get("activity_attempt_id") or task.get("attempt_id", "") + try: + status = await asyncio.wait_for(self.client.activity_task_status( + task_id=task_id, activity_attempt_id=attempt_id, lease_owner=self.worker_id, + ), timeout=5.0) + receipt = status.get("cancellation_acknowledgement") if isinstance(status, Mapping) else None + if (not isinstance(status, Mapping) or status.get("task_id") != task_id + or status.get("activity_attempt_id") != attempt_id or status.get("lease_owner") != self.worker_id + or status.get("can_continue") is not False or status.get("cancel_requested") is not True + or status.get("heartbeat_recorded") is not False or not isinstance(receipt, Mapping) + or receipt.get("callback_state") not in ("unknown", "stopped")): + return + for field in ("request_id", "root_request_id", "cleanup_deadline_at", "cancellation_history_event_id"): + if not isinstance(receipt.get(field), str) or not receipt[field].strip(): + return + reply = await asyncio.wait_for(self.client.acknowledge_activity_cancellation( + task_id=task_id, activity_attempt_id=attempt_id, lease_owner=self.worker_id, + request_id=receipt["request_id"], + ), timeout=5.0) + if (not isinstance(reply, Mapping) or reply.get("task_id") != task_id + or reply.get("activity_attempt_id") != attempt_id or reply.get("lease_owner") != self.worker_id + or reply.get("request_id") != receipt["request_id"] or reply.get("acknowledged") is not True + or not isinstance(reply.get("duplicate"), bool) or reply.get("reason") is not None + or reply.get("heartbeat_recorded") is not False + or not isinstance(reply.get("history_event_id"), str) or not reply["history_event_id"].strip()): + raise _RemoteActivityExecutionAborted("Server did not prove the original callback-stop receipt") + log.info("remote activity %s callback stopped for request %s, root %s, deadline %s", + task_id, receipt["request_id"], receipt["root_request_id"], receipt["cleanup_deadline_at"]) + except Exception as error: + log.warning("remote activity %s stopped but its cancellation receipt failed: %s", task_id, error) + async def _assert_remote_activity_claim(self, task: dict[str, Any]) -> None: try: if self._local_activity_shutdown.is_set(): @@ -2262,24 +2319,8 @@ async def _assert_remote_activity_claim(self, task: dict[str, Any]) -> None: async def _execute_cooperative_remote_callable( self, task: dict[str, Any], handler: Callable[..., Any], args: tuple[Any, ...], info: ActivityInfo, ) -> Any: - abandoned = threading.Event() - owner_loop = asyncio.get_running_loop() - - def boundary() -> None: - if abandoned.is_set() or self._local_activity_shutdown.is_set(): - raise _RemoteActivityExecutionAborted("remote activity callback no longer owns its claim") - - async def observe() -> None: - boundary() - try: - await self._assert_remote_activity_claim(task) - except BaseException: - abandoned.set() - raise - boundary() - async def send_heartbeat(details: dict[str, Any] | None) -> None: - await observe() + await self._assert_remote_activity_claim(task) try: reply = await asyncio.wait_for(self.client.heartbeat_activity_task( task_id=info.task_id, activity_attempt_id=info.activity_attempt_id, @@ -2291,50 +2332,24 @@ async def send_heartbeat(details: dict[str, Any] | None) -> None: or reply.get("cancel_requested") is not False): raise _RemoteActivityExecutionAborted("remote activity heartbeat lost its ownership fence") except Exception as error: - abandoned.set() raise _RemoteActivityExecutionAborted("remote activity user heartbeat failed") from error - await observe() - - async def heartbeat(details: dict[str, Any] | None) -> None: - boundary() - if asyncio.get_running_loop() is owner_loop: - await send_heartbeat(details) - else: - # A synchronous handler may use its own loop to await heartbeat. - # Its HTTP client and claim checks remain on the owner loop. - proxy = asyncio.run_coroutine_threadsafe(send_heartbeat(details), owner_loop) - await asyncio.wrap_future(proxy) - - if inspect.iscoroutinefunction(handler): - async def guarded(*arguments: Any) -> Any: - boundary() - return await handler(*arguments) - callback: Callable[..., Any] = guarded - else: - def guarded_sync(*arguments: Any) -> Any: - boundary() - return handler(*arguments) - callback = guarded_sync - - async def invoke() -> Any: - _set_context(ActivityContext(info=info, client=self.client, heartbeat_callback=heartbeat)) - try: - return await self._execute_activity_callable( - task, info.activity_type, args, callback, run_sync_in_thread=True, remote=True, - ) - finally: - _set_context(None) + await self._assert_remote_activity_claim(task) async def observe_ownership() -> None: while True: await asyncio.sleep(1.0) - await observe() + await self._assert_remote_activity_claim(task) - await observe() - invocation = asyncio.create_task(invoke()) + await self._assert_remote_activity_claim(task) + callback = SupervisedCallback(CallbackInvocation(handler, args, info, task, self.interceptors)) + self._remote_activity_processes.add(callback) + invocation: asyncio.Task[Any] | None = None + abandoned = False observation = asyncio.create_task(observe_ownership()) shutdown = asyncio.create_task(self._local_activity_shutdown.wait()) try: + await callback.start() + invocation = asyncio.create_task(callback.result(send_heartbeat)) done, _ = await asyncio.wait([invocation, observation, shutdown], return_when=asyncio.FIRST_COMPLETED) if shutdown in done: raise _RemoteActivityExecutionAborted("worker shutdown abandoned its remote activity claim") @@ -2342,25 +2357,33 @@ async def observe_ownership() -> None: await observation raise _RemoteActivityExecutionAborted("remote ownership observer stopped without a response") try: - result = await invocation - except Exception: - await observe() # A genuine application failure also needs current authority. - raise - await observe() # Fence before the Client encodes or externalizes the result. - return result + outcome = await invocation + except CallbackProcessLost as error: + raise _RemoteActivityExecutionAborted("remote callback supervisor lost stop authority") from error + await self._assert_remote_activity_claim(task) + if outcome.failure is not None: + raise _RemoteActivityApplicationFailure(outcome.failure) + return outcome.value + except (CallbackProcessLost, _RemoteActivityExecutionAborted, asyncio.CancelledError): + abandoned = True + raise finally: - abandoned.set() - - def discard_late_result(future: asyncio.Task[Any]) -> None: - if not future.cancelled(): - future.exception() - - # The abandoned claim must finish without waiting for callback or - # observer cancellation. All late progress/publication stays fenced. - for background in (invocation, observation, shutdown): + backgrounds = [background for background in (invocation, observation, shutdown) if background is not None] + for background in backgrounds: if not background.done(): background.cancel() - background.add_done_callback(discard_late_result) + await asyncio.gather(*backgrounds, return_exceptions=True) + try: + if not callback.stopped: + await callback.stop() + except Exception as error: + log.warning("remote activity %s callback stop remains unconfirmed: %s", info.task_id, error) + if callback.stopped: + self._remote_activity_processes.discard(callback) + if abandoned: + await self._acknowledge_stopped_remote_activity(task) + # An unconfirmed callback retains its capacity and never emits a + # stop receipt, result or failure. Owner shutdown closes its channel. async def _run_activity_task(self, task: dict[str, Any]) -> str: self._track_worker_session_from_task(task) @@ -2455,6 +2478,23 @@ async def _run_activity_task(self, task: dict[str, Any]) -> str: except _RemoteActivityExecutionAborted as error: log.warning("remote activity %s claim abandoned: %s", task_id, error) return "claim_aborted" + except CallbackProcessLost as error: + log.warning("remote activity %s supervision failed: %s", task_id, error) + return "claim_aborted" + except _RemoteActivityApplicationFailure as error: + failure = error.failure + try: + await self.client.fail_activity_task( + task_id=task_id, activity_attempt_id=attempt_id, lease_owner=self.worker_id, + message="activity cancelled" if failure.cancelled else failure.message, + failure_type=failure.failure_type, failure_class=failure.failure_class, + failure_code=failure.failure_code, stack_trace=failure.stack_trace, + non_retryable=failure.non_retryable or failure.cancelled, + ) + except Exception as report_error: + log.warning("failed to report remote activity failure: %s", report_error) + return ("cancelled" if failure.cancelled else + "failed_non_retryable" if failure.non_retryable else "failed") except ActivityCancelled: log.info("activity %s cancelled via heartbeat", task_id) try: @@ -3027,7 +3067,9 @@ async def _report_unhandled_workflow_task_error( async def _poll_activity_tasks(self) -> None: while not self._stop.is_set(): - if len(self._remote_activity_threads) >= self.max_concurrent_activity_tasks: + self._validate_cooperative_activity_handlers() + if (len(self._remote_activity_threads) + len(self._remote_activity_processes) + >= self.max_concurrent_activity_tasks): await asyncio.sleep(0.1) continue await self._act_semaphore.acquire() @@ -3465,7 +3507,8 @@ def _current_task_slots(self) -> dict[str, int]: ), "activity_available": max( 0, min(self.max_concurrent_activity_tasks - self._activity_inflight, - self.max_concurrent_activity_tasks - len(self._remote_activity_threads)), + self.max_concurrent_activity_tasks - len(self._remote_activity_threads) + - len(self._remote_activity_processes)), ), "session_available": max( 0, self.max_concurrent_worker_sessions @@ -3669,6 +3712,11 @@ async def _run_until_loop( next_task_kind = "workflow" continue + self._validate_cooperative_activity_handlers() + if len(self._remote_activity_processes) >= self.max_concurrent_activity_tasks: + next_task_kind = "workflow" + await asyncio.sleep(poll_interval) + continue poll_start = time.perf_counter() task = await self.client.poll_activity_task( worker_id=self.worker_id, @@ -3786,6 +3834,11 @@ async def _shutdown(self) -> None: log.warning("cancelled %d task(s) after shutdown timeout", len(pending)) await asyncio.sleep(0) still_running = {task for task in pending if not task.done()} + if still_running and self._remote_activity_processes: + # Application drain has ended. Allow bounded process reaping and + # stop-report transport without granting more callback work or + # extending the run's original cancellation deadline. + _, still_running = await asyncio.wait(still_running, timeout=15.0) if still_running: raise RuntimeError( "worker shutdown timed out while cancelling " @@ -3799,6 +3852,12 @@ async def _shutdown(self) -> None: if self._remote_activity_executor is not None: self._remote_activity_executor.shutdown(wait=False, cancel_futures=True) + if self._remote_activity_processes: + raise RuntimeError( + "worker shutdown has unconfirmed remote callback stop(s); " + "the worker registration remains active" + ) + for session in self._worker_sessions.values(): if not session.active: continue diff --git a/tests/test_cooperative_cancellation_worker.py b/tests/test_cooperative_cancellation_worker.py index fb978fa..692a337 100644 --- a/tests/test_cooperative_cancellation_worker.py +++ b/tests/test_cooperative_cancellation_worker.py @@ -135,12 +135,16 @@ async def complete(self, **kwargs: Any) -> dict[str, Any]: async def worker(self, monkeypatch: pytest.MonkeyPatch, **kwargs: Any) -> Worker: monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + # These transport/replay fixtures use local callbacks that share mock + # state with this process. They never poll or advertise remote handlers. + local_callbacks = kwargs.pop("activities", []) worker = Worker( self.client, task_queue="queue", worker_id=self.lease_owner, workflows=kwargs.pop("workflows", [CancellationWorkflow]), capabilities=["cooperative_cancellation"], **kwargs, ) await worker._register() + worker.activities.update({callback.__activity_name__: callback for callback in local_callbacks}) return worker diff --git a/tests/test_cooperative_remote_worker.py b/tests/test_cooperative_remote_worker.py index c9efe27..ca61e36 100644 --- a/tests/test_cooperative_remote_worker.py +++ b/tests/test_cooperative_remote_worker.py @@ -1,8 +1,11 @@ from __future__ import annotations import asyncio -import threading +import os +import shutil +import time from copy import deepcopy +from pathlib import Path from typing import Any from unittest.mock import AsyncMock, patch @@ -12,15 +15,16 @@ from durable_workflow import activity, serializer from durable_workflow.client import Client -from durable_workflow.errors import ActivityCancelled, NonRetryableError +from durable_workflow.errors import ActivityCancelled, NonRetryableError, ServerError from durable_workflow.retry_policy import TransportRetryPolicy -from durable_workflow.worker import Worker, _RemoteActivityExecutionAborted +from durable_workflow.worker import Worker +from tests.test_activity_process import wait_for_exit, wait_for_file from tests.test_worker import compatible_cluster_info -def task() -> dict[str, Any]: +def task(*args: Any) -> dict[str, Any]: return {"task_id": "remote-task", "activity_attempt_id": "attempt", "activity_type": "remote", - "payload_codec": "avro", "arguments": serializer.envelope([], codec="avro")} + "payload_codec": "avro", "arguments": serializer.envelope(list(args), codec="avro")} def status() -> dict[str, Any]: @@ -29,6 +33,58 @@ def status() -> dict[str, Any]: "lease_expires_at": "2100-01-01T00:00:00Z", "deadlines": None, "worker_session": None} +def cancellation_status() -> dict[str, Any]: + return {**status(), "can_continue": False, "cancel_requested": True, "reason": "activity_cancelled", + "cancellation_acknowledgement": {"request_id": "original-request", "root_request_id": "root-request", + "cleanup_deadline_at": "2100-01-01T00:00:00Z", + "cancellation_history_event_id": "cancellation-event", + "callback_state": "unknown"}} + + +def stop_receipt() -> dict[str, Any]: + return {"task_id": "remote-task", "activity_attempt_id": "attempt", "lease_owner": "owner", + "request_id": "original-request", "acknowledged": True, "duplicate": False, "reason": None, + "heartbeat_recorded": False, "history_event_id": "receipt-event"} + + +async def blocked_async(marker: str, heartbeat: bool = False) -> None: + if heartbeat: + await activity.context().heartbeat({"authored": True}) + Path(marker).write_text(str(os.getpid())) + await asyncio.Event().wait() + + +def blocked_sync(marker: str, heartbeat: bool = False) -> None: + if heartbeat: + asyncio.run(activity.context().heartbeat({"authored": True})) + Path(marker).write_text(str(os.getpid())) + time.sleep(60) + + +def released_sync(marker: str) -> bytes: + Path(marker).write_text(str(os.getpid())) + while not Path(marker + ".release").exists(): + time.sleep(0.02) + return b"late result" + + +async def typed_async() -> dict[str, Any]: + await activity.context().heartbeat({"authored": True}) + return {"bytes": b"\x00\xff", "value": 42, "null": None, "task_id": activity.context().info.task_id} + + +def typed_sync() -> dict[str, Any]: + return asyncio.run(typed_async()) + + +async def application_failure(kind: str) -> None: + if kind == "ValueError": + raise ValueError("original") + if kind == "NonRetryableError": + raise NonRetryableError("original") + raise ActivityCancelled() + + @pytest_asyncio.fixture async def owner(monkeypatch: pytest.MonkeyPatch): # type: ignore[no-untyped-def] monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") @@ -39,6 +95,7 @@ async def owner(monkeypatch: pytest.MonkeyPatch): # type: ignore[no-untyped-def client.register_worker.return_value = {"registered": True} client.activity_task_status.return_value = status() client.heartbeat_activity_task.return_value = status() + client.acknowledge_activity_cancellation.return_value = stop_receipt() worker = Worker(client, task_queue="queue", worker_id="owner", capabilities=["cooperative_cancellation"], max_concurrent_activity_tasks=1) await worker._register() @@ -49,106 +106,63 @@ async def owner(monkeypatch: pytest.MonkeyPatch): # type: ignore[no-untyped-def @pytest.mark.parametrize("reason", ["cancel", "replacement", "backend", "shutdown"]) -async def test_blocked_async_callback_is_abandoned_without_progress_or_publication(owner, reason: str) -> None: +@pytest.mark.parametrize("synchronous", [False, True]) +@pytest.mark.parametrize("heartbeat", [False, True]) +async def test_authority_loss_stops_and_joins_the_actual_callback_before_reporting( + owner: Any, tmp_path: Path, reason: str, synchronous: bool, heartbeat: bool, +) -> None: worker, client = owner - entered, late_fenced = asyncio.Event(), asyncio.Event() - - async def callback() -> object: - entered.set() - try: - await asyncio.Event().wait() - except asyncio.CancelledError: - with pytest.raises(_RemoteActivityExecutionAborted): - await activity.context().heartbeat({"late": True}) - late_fenced.set() - return object() - - worker.activities["remote"] = callback - execution = asyncio.create_task(worker._run_activity_task(task())) - await asyncio.wait_for(entered.wait(), timeout=2) + marker = tmp_path / "callback" + worker.activities["remote"] = blocked_sync if synchronous else blocked_async + execution = asyncio.create_task(worker._run_activity_task(task(str(marker), heartbeat))) + await wait_for_file(marker) + pid = int(marker.read_text()) + + async def acknowledge(**options: Any) -> dict[str, Any]: + with pytest.raises(ProcessLookupError): + os.kill(pid, 0) + assert not worker._remote_activity_processes + assert options == {"task_id": "remote-task", "activity_attempt_id": "attempt", "lease_owner": "owner", + "request_id": "original-request"} + return stop_receipt() + + client.acknowledge_activity_cancellation.side_effect = acknowledge if reason == "cancel": - client.activity_task_status.return_value = {**status(), "can_continue": False, - "cancel_requested": True, "reason": "activity_cancelled"} + client.activity_task_status.return_value = cancellation_status() elif reason == "replacement": client.activity_task_status.return_value = {**status(), "activity_attempt_id": "replacement"} elif reason == "backend": client.activity_task_status.side_effect = OSError("backend unavailable") else: worker._local_activity_shutdown.set() - assert await asyncio.wait_for(execution, timeout=3) == "claim_aborted" - await asyncio.wait_for(late_fenced.wait(), timeout=2) - client.heartbeat_activity_task.assert_not_awaited() + assert await asyncio.wait_for(execution, timeout=4.0) == "claim_aborted" + await wait_for_exit(pid) + assert not worker._remote_activity_processes + assert worker._current_task_slots()["activity_available"] == 1 + assert client.heartbeat_activity_task.await_count == int(heartbeat) + assert client.acknowledge_activity_cancellation.await_count == int(reason == "cancel") client.complete_activity_task.assert_not_awaited() client.fail_activity_task.assert_not_awaited() -async def test_synchronous_callback_keeps_owner_available_and_fences_late_thread_heartbeat(owner) -> None: - worker, client = owner - entered, late_fenced = asyncio.Event(), asyncio.Event() - release = threading.Event() - loop = asyncio.get_running_loop() - - def callback() -> object: - loop.call_soon_threadsafe(entered.set) - assert release.wait(timeout=10) - with pytest.raises(_RemoteActivityExecutionAborted): - asyncio.run(activity.context().heartbeat({"late": True})) - loop.call_soon_threadsafe(late_fenced.set) - return object() - - worker.activities["remote"] = callback - execution = asyncio.create_task(worker._run_activity_task(task())) - try: - await asyncio.wait_for(entered.wait(), timeout=2) - client.activity_task_status.return_value = {**status(), "can_continue": False, "reason": "lease_expired"} - assert await asyncio.wait_for(execution, timeout=3) == "claim_aborted" - assert not late_fenced.is_set() # The Python thread is still running, with its attempt fenced. - assert worker._current_task_slots()["activity_available"] == 0 - poller = asyncio.create_task(worker._poll_activity_tasks()) - try: - await asyncio.sleep(0.05) - client.poll_activity_task.assert_not_awaited() - finally: - poller.cancel() - with pytest.raises(asyncio.CancelledError): - await poller - release.set() - await asyncio.wait_for(late_fenced.wait(), timeout=2) - client.heartbeat_activity_task.assert_not_awaited() - client.complete_activity_task.assert_not_awaited() - client.fail_activity_task.assert_not_awaited() - finally: - release.set() - - -async def test_shutdown_expiry_of_tracked_remote_attempt_fences_a_cancellation_resistant_result(owner) -> None: +async def test_shutdown_timeout_joins_the_remote_callback_before_returning(owner: Any, tmp_path: Path) -> None: worker, client = owner worker._shutdown_timeout = 0.01 - entered, late_fenced = asyncio.Event(), asyncio.Event() - - async def callback() -> object: - entered.set() - try: - await asyncio.Event().wait() - except asyncio.CancelledError: - with pytest.raises(_RemoteActivityExecutionAborted): - await activity.context().heartbeat() - late_fenced.set() - return object() - - worker.activities["remote"] = callback - execution = worker._track(worker._run_activity_task(task())) - await asyncio.wait_for(entered.wait(), timeout=2) - await asyncio.wait_for(worker.stop(), timeout=2) - await asyncio.wait_for(late_fenced.wait(), timeout=2) + marker = tmp_path / "callback" + worker.activities["remote"] = blocked_async + execution = worker._track(worker._run_activity_task(task(str(marker)))) + await wait_for_file(marker) + pid = int(marker.read_text()) + await asyncio.wait_for(worker.stop(), timeout=4.0) + await wait_for_exit(pid) assert execution.cancelled() or execution.result() == "claim_aborted" - client.heartbeat_activity_task.assert_not_awaited() + assert not worker._remote_activity_processes client.complete_activity_task.assert_not_awaited() client.fail_activity_task.assert_not_awaited() @pytest.mark.parametrize("synchronous", [False, True]) -async def test_authored_heartbeat_stays_on_owner_loop_and_preserves_typed_result(owner, synchronous: bool) -> None: +async def test_authored_heartbeat_stays_on_owner_loop_and_preserves_typed_result(owner: Any, synchronous: bool) -> None: worker, client = owner loop = asyncio.get_running_loop() observed_loops = [] @@ -158,26 +172,20 @@ async def heartbeat(**kwargs: Any) -> dict[str, Any]: assert kwargs["details"] == {"authored": True} return status() - async def asynchronous() -> dict[str, Any]: - await activity.context().heartbeat({"authored": True}) - return {"bytes": b"\x00\xff", "value": 42, "null": None} - - def blocking() -> dict[str, Any]: - asyncio.run(activity.context().heartbeat({"authored": True})) - return {"bytes": b"\x00\xff", "value": 42, "null": None} - client.heartbeat_activity_task.side_effect = heartbeat - worker.activities["remote"] = blocking if synchronous else asynchronous + worker.activities["remote"] = typed_sync if synchronous else typed_async assert await worker._run_activity_task(task()) == "completed" assert observed_loops == [loop] assert client.complete_activity_task.await_args.kwargs["result"] == { - "bytes": b"\x00\xff", "value": 42, "null": None, + "bytes": b"\x00\xff", "value": 42, "null": None, "task_id": "remote-task", } + assert not worker._remote_activity_processes + client.acknowledge_activity_cancellation.assert_not_awaited() client.fail_activity_task.assert_not_awaited() @pytest.mark.parametrize("field", ["lease", "heartbeat", "start_to_close", "schedule_to_close", "session", "malformed"]) -async def test_invalid_or_elapsed_observed_bounds_prevent_callback_start(owner, field: str) -> None: +async def test_invalid_or_elapsed_observed_bounds_prevent_callback_start(owner: Any, field: str) -> None: worker, client = owner reply = deepcopy(status()) expired = "2000-01-01T00:00:00Z" @@ -195,25 +203,96 @@ async def test_invalid_or_elapsed_observed_bounds_prevent_callback_start(owner, worker.activities["remote"] = callback assert await worker._run_activity_task(task()) == "claim_aborted" callback.assert_not_awaited() + client.acknowledge_activity_cancellation.assert_not_awaited() client.complete_activity_task.assert_not_awaited() client.fail_activity_task.assert_not_awaited() -@pytest.mark.parametrize("error", [ValueError("original"), NonRetryableError("original"), ActivityCancelled()]) -async def test_genuine_application_failure_preserves_classification(owner, error: Exception) -> None: +@pytest.mark.parametrize("kind", ["ValueError", "NonRetryableError", "ActivityCancelled"]) +async def test_genuine_application_failure_preserves_classification(owner: Any, kind: str) -> None: worker, client = owner + worker.activities["remote"] = application_failure + outcome = await worker._run_activity_task(task(kind)) + assert outcome in {"failed", "failed_non_retryable", "cancelled"} + report = client.fail_activity_task.await_args.kwargs + assert report["failure_type"] == kind + assert report["non_retryable"] is (kind != "ValueError") + assert report["failure_class"].endswith("." + kind) + assert "application_failure" in report["stack_trace"] + client.complete_activity_task.assert_not_awaited() - async def callback() -> None: - raise error - worker.activities["remote"] = callback - outcome = await worker._run_activity_task(task()) - assert outcome in {"failed", "failed_non_retryable", "cancelled"} - assert client.fail_activity_task.await_args.kwargs["failure_type"] == type(error).__name__ - assert client.fail_activity_task.await_args.kwargs.get("non_retryable", False) is isinstance( - error, (NonRetryableError, ActivityCancelled), - ) +@pytest.mark.parametrize("refusal", ["refused", "unproved"]) +async def test_stop_receipt_failure_never_becomes_result_or_failure(owner: Any, tmp_path: Path, refusal: str) -> None: + worker, client = owner + marker = tmp_path / "callback" + worker.activities["remote"] = blocked_async + execution = asyncio.create_task(worker._run_activity_task(task(str(marker)))) + await wait_for_file(marker) + if refusal == "refused": + client.acknowledge_activity_cancellation.side_effect = ServerError(409, {"reason": "request_mismatch"}) + else: + client.acknowledge_activity_cancellation.return_value = {**stop_receipt(), "acknowledged": False} + client.activity_task_status.return_value = cancellation_status() + assert await asyncio.wait_for(execution, timeout=4.0) == "claim_aborted" + await wait_for_exit(int(marker.read_text())) + client.acknowledge_activity_cancellation.assert_awaited_once() client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + + +async def test_dead_supervisor_retains_capacity_and_never_reports_stop(owner: Any, tmp_path: Path) -> None: + worker, client = owner + marker = tmp_path / "callback" + worker.activities["remote"] = released_sync + execution = asyncio.create_task(worker._run_activity_task(task(str(marker)))) + await wait_for_file(marker) + callback = next(iter(worker._remote_activity_processes)) + pid = int(marker.read_text()) + try: + callback._supervisor.kill() + assert await asyncio.wait_for(execution, timeout=4.0) == "claim_aborted" + os.kill(pid, 0) + assert callback.stopped is False + assert worker._current_task_slots()["activity_available"] == 0 + poller = asyncio.create_task(worker._poll_activity_tasks()) + try: + await asyncio.sleep(0.1) + client.poll_activity_task.assert_not_awaited() + finally: + poller.cancel() + with pytest.raises(asyncio.CancelledError): + await poller + client.acknowledge_activity_cancellation.assert_not_awaited() + client.complete_activity_task.assert_not_awaited() + client.fail_activity_task.assert_not_awaited() + with pytest.raises(RuntimeError, match="unconfirmed remote callback stop.*registration remains active"): + await worker._shutdown() + assert worker._registered is True + client.deregister_worker_registration.assert_not_awaited() + finally: + Path(str(marker) + ".release").touch() + await wait_for_exit(pid) + shutil.rmtree(callback.directory, ignore_errors=True) + # The test performs external recovery after proving the worker cannot + # confirm stop. Release this synthetic state only after the process dies. + worker._remote_activity_processes.discard(callback) + + +async def test_registration_refuses_nonimportable_cooperative_handler_before_claiming(owner: Any) -> None: + _, client = owner + + @activity.defn(name="not-importable") + async def callback() -> None: + pass + + client.register_worker.reset_mock() + worker = Worker(client, task_queue="queue", worker_id="unsupported-owner", activities=[callback], + capabilities=["cooperative_cancellation"]) + with pytest.raises(RuntimeError, match="unsupported-owner.*not-importable.*spawn-compatible"): + await worker._register() + client.register_worker.assert_not_awaited() + client.poll_activity_task.assert_not_awaited() async def test_client_observation_keeps_worker_credentials_namespace_and_fence(monkeypatch: pytest.MonkeyPatch) -> None: From edf890a9d160fded76ce1e3e7d15bac8a997f4f5 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 02:01:02 +0000 Subject: [PATCH 21/41] Prove connected Python callback stops and original cancellation receipts --- tests/integration/cooperative_worker.py | 32 +-- .../test_cooperative_cancellation.py | 215 +++++++++++------- 2 files changed, 145 insertions(+), 102 deletions(-) diff --git a/tests/integration/cooperative_worker.py b/tests/integration/cooperative_worker.py index 3b95e26..63f84ab 100644 --- a/tests/integration/cooperative_worker.py +++ b/tests/integration/cooperative_worker.py @@ -11,32 +11,36 @@ from tests.integration.test_cooperative_cancellation import candidate_worker, cooperative_cleanup, poll_claim +async def cleanup(request_id: str) -> str: + info = activity.context().info + print(json.dumps({"phase": "cleanup", "request_id": request_id, "worker_id": info.worker_id}), flush=True) + if os.environ.get("DW_COOPERATIVE_FIXTURE_MODE") == "hold": + await asyncio.Event().wait() + return await cooperative_cleanup(request_id) + + +async def blocked_remote() -> None: + info = activity.context().info + print(json.dumps({"phase": "remote-entered", "task_id": info.task_id, + "activity_attempt_id": info.activity_attempt_id, + "lease_owner": info.worker_id, "attempt_number": info.attempt_number, + "callback_pid": os.getpid()}), flush=True) + await asyncio.Event().wait() + + async def main() -> None: queue, worker_id, mode = sys.argv[1:] if mode not in {"hold", "finish", "remote"}: raise ValueError("unknown qualification mode") + os.environ["DW_COOPERATIVE_FIXTURE_MODE"] = mode async with Client( os.environ["DURABLE_WORKFLOW_SERVER_URL"], token=os.environ.get("DURABLE_WORKFLOW_AUTH_TOKEN", "test-token"), namespace="default", ) as client: worker = candidate_worker(client, queue, worker_id=worker_id) - async def cleanup(request_id: str) -> str: - print(json.dumps({"phase": "cleanup", "request_id": request_id, "worker_id": worker_id}), flush=True) - if mode == "hold": - await asyncio.Event().wait() - return await cooperative_cleanup(request_id) - worker.activities["tests.python-cooperative-cleanup"] = cleanup if mode == "remote": - async def blocked_remote() -> object: - info = activity.context().info - print(json.dumps({"phase": "remote-entered", "task_id": info.task_id, - "activity_attempt_id": info.activity_attempt_id, - "lease_owner": info.worker_id, - "attempt_number": info.attempt_number}), flush=True) - await asyncio.Event().wait() - return object() worker.activities["tests.python-cooperative-work"] = blocked_remote await worker.run() return diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 6a10040..7ddf4ef 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -9,7 +9,9 @@ import threading import time import uuid +from dataclasses import dataclass from datetime import datetime +from pathlib import Path from typing import Any import pytest @@ -17,7 +19,7 @@ from durable_workflow import Client, Worker, activity, workflow from durable_workflow.client import WorkflowHandle from durable_workflow.errors import ServerError, WorkflowCancelled -from durable_workflow.worker import _poll_capacity_delay, _RemoteActivityExecutionAborted +from durable_workflow.worker import _poll_capacity_delay from durable_workflow.workflow import LocalActivityExecutionAborted pytestmark = pytest.mark.usefixtures("cooperative_runtime") @@ -135,137 +137,165 @@ async def heartbeat_worker(self, **kwargs: Any) -> Any: return reply -@pytest.mark.parametrize("handler_kind", ["async", "sync"]) -@pytest.mark.parametrize("user_heartbeat", [False, True]) -async def test_actual_remote_worker_fences_blocked_callbacks_without_manufacturing_progress( - server_url: str, server_token: str, handler_kind: str, user_heartbeat: bool, -) -> None: - queue = f"py-cooperative-owner-{uuid.uuid4().hex[:8]}" - entered, late_fenced = asyncio.Event(), asyncio.Event() - release_thread = threading.Event() - loop = asyncio.get_running_loop() - fences: list[dict[str, str]] = [] +@dataclass(frozen=True) +class AsyncRemoteQualification: + marker: str + user_heartbeat: bool = False - def record_fence() -> None: - info = activity.context().info - fences.append({"task_id": info.task_id, "activity_attempt_id": info.activity_attempt_id, - "lease_owner": info.worker_id}) + async def __call__(self) -> None: + if self.user_heartbeat: + await activity.context().heartbeat({"qualification": "remote-in-flight"}) + write_remote_marker(self.marker) + await asyncio.Event().wait() - async def asynchronous() -> object: - record_fence() - entered.set() - try: - while True: - await asyncio.sleep(0.1) - if user_heartbeat: - await activity.context().heartbeat({"qualification": "remote-in-flight"}) - except (asyncio.CancelledError, _RemoteActivityExecutionAborted): - with pytest.raises(_RemoteActivityExecutionAborted): - await activity.context().heartbeat({"late": True}) - late_fenced.set() - return object() - def synchronous() -> object: - record_fence() - loop.call_soon_threadsafe(entered.set) - try: - while not release_thread.wait(timeout=0.1): - if user_heartbeat: - asyncio.run(activity.context().heartbeat({"qualification": "remote-in-flight"})) - except _RemoteActivityExecutionAborted: - pass - with pytest.raises(_RemoteActivityExecutionAborted): - asyncio.run(activity.context().heartbeat({"late": True})) - loop.call_soon_threadsafe(late_fenced.set) - return object() +@dataclass(frozen=True) +class SyncRemoteQualification: + marker: str + user_heartbeat: bool = False + + def __call__(self) -> None: + if self.user_heartbeat: + asyncio.run(activity.context().heartbeat({"qualification": "remote-in-flight"})) + write_remote_marker(self.marker) + while True: + time.sleep(1) + + +def write_remote_marker(marker: str) -> None: + info = activity.context().info + pending = Path(marker + ".writing") + pending.write_text(json.dumps({ + "task_id": info.task_id, "activity_attempt_id": info.activity_attempt_id, + "lease_owner": info.worker_id, "callback_pid": os.getpid(), + })) + pending.replace(marker) + +async def remote_marker(marker: Path) -> dict[str, Any]: + async def ready() -> dict[str, Any]: + while not marker.exists(): + await asyncio.sleep(0.05) + return json.loads(marker.read_text()) + return await asyncio.wait_for(ready(), timeout=20) + + +async def callback_gone(pid: int) -> None: + async def gone() -> None: + while True: + try: + os.kill(pid, 0) + except ProcessLookupError: + return + await asyncio.sleep(0.05) + await asyncio.wait_for(gone(), timeout=10) + + +async def observed_stop_receipt(client: Client, fence: dict[str, str]) -> dict[str, Any]: + async def receipt() -> dict[str, Any]: + while True: + reply = await client.activity_task_status(**fence) + proof = reply.get("cancellation_acknowledgement") + if isinstance(proof, dict) and proof.get("callback_state") == "stopped": + assert reply["heartbeat_recorded"] is False + assert reply["can_continue"] is False + return proof + await asyncio.sleep(0.1) + return await asyncio.wait_for(receipt(), timeout=15) + + +@pytest.mark.parametrize("handler_kind", ["async", "sync"]) +@pytest.mark.parametrize("user_heartbeat", [False, True]) +async def test_actual_remote_worker_stops_callbacks_and_reports_original_cancellation( + server_url: str, server_token: str, handler_kind: str, user_heartbeat: bool, tmp_path: Path, +) -> None: + queue = f"py-cooperative-owner-{uuid.uuid4().hex[:8]}" + marker = tmp_path / "remote" async with ObservedOwnerClient(server_url, token=server_token, namespace="default") as client: worker = candidate_worker(client, queue, max_concurrent_activity_tasks=1) - worker.activities["tests.python-cooperative-work"] = asynchronous if handler_kind == "async" else synchronous + worker.activities["tests.python-cooperative-work"] = ( + AsyncRemoteQualification(str(marker), user_heartbeat) if handler_kind == "async" + else SyncRemoteQualification(str(marker), user_heartbeat) + ) running = asyncio.create_task(worker.run()) try: handle = await client.start_workflow( workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], ) - await asyncio.wait_for(entered.wait(), timeout=15) + entered = await remote_marker(marker) await asyncio.wait_for(client.owner_heartbeat.wait(), timeout=15) - assert not late_fenced.is_set() + fence = {key: entered[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) - history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) - release_thread.set() - await asyncio.wait_for(late_fenced.wait(), timeout=5) + original = accepted["cancellation_request"] + await assert_cancelled_cleanup(handle, original["request_id"]) + await callback_gone(entered["callback_pid"]) + if os.environ.get("DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION") == "1": + proof = await observed_stop_receipt(client, fence) + assert proof["request_id"] == original["request_id"] + assert proof["root_request_id"] == original["request_id"] + assert proof["cleanup_deadline_at"] == original["cleanup_deadline_at"] + assert proof["received_after_deadline"] is False + duplicate = await client.acknowledge_activity_cancellation(**fence, request_id=original["request_id"]) + assert duplicate["duplicate"] is True + assert duplicate["history_event_id"] == proof["history_event_id"] + print(f"Remote callback stop receipt: {json.dumps([entered, original, proof, duplicate])}") + history = await events(handle) + if os.environ.get("DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION") == "1": + receipt = [event for event in history if event["event_type"] == "ActivityCancellationAcknowledged"] + assert len(receipt) == 1 + assert receipt[0]["id"] == proof["history_event_id"] + assert receipt[0]["payload"]["evidence_source"] == "activity_worker" delivery = [event for event in history if event["event_type"] == "CooperativeCancellationDelivered"][0] assert delivery["payload"]["call_kind"] == "activity" assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 progress = [event for event in history if event["event_type"] == "ActivityHeartbeatRecorded" and event["payload"].get("activity_type") == "tests.python-cooperative-work"] assert bool(progress) is user_heartbeat - assert len(fences) == 1 with pytest.raises(ServerError) as completion: - await client.complete_activity_task(**fences[0], result="late") + await client.complete_activity_task(**fence, result="late") assert completion.value.status == 409 with pytest.raises(ServerError) as failure: - await client.fail_activity_task(**fences[0], message="late", failure_type="LateQualification") + await client.fail_activity_task(**fence, message="late", failure_type="LateQualification") assert failure.value.status == 409 assert await events(handle) == history + print(f"Remote stopped cancellation history: {json.dumps(history)}") finally: - release_thread.set() await worker.stop() - await asyncio.wait_for(running, timeout=5) + await asyncio.wait_for(running, timeout=10) @pytest.mark.parametrize("handler_kind", ["async", "sync"]) -async def test_actual_remote_shutdown_expiry_fences_callback_before_replacement_cleanup( - server_url: str, server_token: str, handler_kind: str, +async def test_actual_remote_shutdown_joins_callback_before_replacement_cleanup( + server_url: str, server_token: str, handler_kind: str, tmp_path: Path, ) -> None: queue = f"py-cooperative-owner-stop-{uuid.uuid4().hex[:8]}" - entered, late_fenced = asyncio.Event(), asyncio.Event() - release = threading.Event() - loop = asyncio.get_running_loop() - - async def asynchronous() -> object: - entered.set() - try: - await asyncio.Event().wait() - except asyncio.CancelledError: - with pytest.raises(_RemoteActivityExecutionAborted): - await activity.context().heartbeat({"late": True}) - late_fenced.set() - return object() - - def synchronous() -> object: - loop.call_soon_threadsafe(entered.set) - assert release.wait(timeout=20) - with pytest.raises(_RemoteActivityExecutionAborted): - asyncio.run(activity.context().heartbeat({"late": True})) - loop.call_soon_threadsafe(late_fenced.set) - return object() - + marker = tmp_path / "remote" async with Client(server_url, token=server_token, namespace="default") as client: worker = candidate_worker(client, queue, shutdown_timeout=0.1) - worker.activities["tests.python-cooperative-work"] = asynchronous if handler_kind == "async" else synchronous + worker.activities["tests.python-cooperative-work"] = ( + AsyncRemoteQualification(str(marker)) if handler_kind == "async" else SyncRemoteQualification(str(marker)) + ) running = asyncio.create_task(worker.run()) replacement = candidate_worker(client, queue, worker_id=f"{queue}-replacement") try: handle = await client.start_workflow( workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], ) - await asyncio.wait_for(entered.wait(), timeout=15) + entered = await remote_marker(marker) before = await events(handle) - await asyncio.wait_for(worker.stop(), timeout=5) - await asyncio.wait_for(running, timeout=5) - release.set() - await asyncio.wait_for(late_fenced.wait(), timeout=5) + await asyncio.wait_for(worker.stop(), timeout=10) + await asyncio.wait_for(running, timeout=10) + await callback_gone(entered["callback_pid"]) assert await events(handle) == before accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) await replacement._register() assert await replacement._run_workflow_task(await poll_claim(client, replacement)) is not None await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + print(f"Stopped remote callback before replacement: {json.dumps(entered)}") finally: - release.set() await replacement.stop() await worker.stop() - await asyncio.wait_for(running, timeout=5) + await asyncio.wait_for(running, timeout=10) async def test_waiting_timer_is_cancelled_by_canonical_delivery( @@ -308,8 +338,10 @@ async def cleanup(request_id: str) -> str: async with Client(server_url, token=server_token, namespace="default") as client: worker = candidate_worker(client, queue) - worker.activities["tests.python-cooperative-cleanup"] = cleanup await worker._register() + # This fixture invokes only local cleanup and claims remote work by API. + # Its in-process closure does not qualify callback process supervision. + worker.activities["tests.python-cooperative-cleanup"] = cleanup try: handle = await client.start_workflow( workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], @@ -411,11 +443,13 @@ async def cleanup(request_id: str) -> str: async with Client(server_url, token=server_token, namespace="default") as client: worker = candidate_worker(client, queue) + await worker._register() + # Exercise the existing local replay fence without polling remote work. + # Physical local callback supervision remains a separate qualification. worker.activities["tests.python-cooperative-work"] = ( blocked_work if handler_kind == "async" else synchronous_work ) worker.activities["tests.python-cooperative-cleanup"] = cleanup - await worker._register() try: handle = await client.start_workflow( workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["local"], @@ -458,8 +492,9 @@ async def cleanup(request_id: str) -> object: async with Client(server_url, token=server_token, namespace="default") as client: worker = candidate_worker(client, queue, shutdown_timeout=5 if shutdown_kind == "drain" else 0.1) - worker.activities["tests.python-cooperative-cleanup"] = cleanup await worker._register() + # Local-only replay fixture. It does not poll remote activity callbacks. + worker.activities["tests.python-cooperative-cleanup"] = cleanup try: handle = await client.start_workflow( workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], @@ -544,6 +579,7 @@ async def test_sigkill_activity_owner_reclaims_attempt_before_cooperative_cleanu print(f"SIGKILL original activity: {json.dumps([original, leased])}") owner.kill() assert await asyncio.wait_for(owner.wait(), timeout=10) == -9 + await callback_gone(original["callback_pid"]) successor = await native_process(queue, f"{queue}-successor", "remote") processes.append(successor) @@ -579,6 +615,7 @@ async def test_sigkill_activity_owner_reclaims_attempt_before_cooperative_cleanu accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + await callback_gone(reclaimed["callback_pid"]) assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 print(f"SIGKILL reclaimed cancellation history: {json.dumps(history)}") finally: @@ -606,6 +643,7 @@ async def test_killed_remote_owner_cannot_publish_after_cold_workflow_delivery( before = await events(handle) owner.kill() assert await asyncio.wait_for(owner.wait(), timeout=10) == -9 + await callback_gone(claimed["callback_pid"]) assert await events(handle) == before accepted = await handle.request_cancellation(cleanup_timeout_seconds=60) replacement = await native_process(queue, f"{queue}-replacement", "finish") @@ -691,8 +729,9 @@ async def cleanup(request_id: str) -> object: async with Client(server_url, token=server_token, namespace="default") as client: worker = candidate_worker(client, queue) - worker.activities["tests.python-cooperative-cleanup"] = cleanup await worker._register() + # Local-only replay fixture. It does not poll remote activity callbacks. + worker.activities["tests.python-cooperative-cleanup"] = cleanup try: handle = await client.start_workflow( workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], From 003ff078a58980898e052da754136131bddb0a8f Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 02:14:28 +0000 Subject: [PATCH 22/41] Compare public cancellation history by immutable receipt payload --- tests/integration/test_cooperative_cancellation.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 7ddf4ef..0bf12b1 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -243,8 +243,14 @@ async def test_actual_remote_worker_stops_callbacks_and_reports_original_cancell if os.environ.get("DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION") == "1": receipt = [event for event in history if event["event_type"] == "ActivityCancellationAcknowledged"] assert len(receipt) == 1 - assert receipt[0]["id"] == proof["history_event_id"] - assert receipt[0]["payload"]["evidence_source"] == "activity_worker" + # Public history exposes event sequence/payload, not row IDs. + # The readonly status and duplicate transport prove receipt ID. + payload = receipt[0]["payload"] + assert payload["activity_attempt_id"] == fence["activity_attempt_id"] + assert payload["lease_owner"] == fence["lease_owner"] + assert payload["evidence_source"] == "activity_worker" + for key in ("request_id", "root_request_id", "cleanup_deadline_at", "acknowledged_at"): + assert payload[key] == proof[key] delivery = [event for event in history if event["event_type"] == "CooperativeCancellationDelivered"][0] assert delivery["payload"]["call_kind"] == "activity" assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 From 8eabc65bb4e01595c3a75758ccca86fb898881d6 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 10:15:54 +0000 Subject: [PATCH 23/41] Add source prepared local activity transport with original claim budgets --- src/durable_workflow/client.py | 92 ++++++++++- tests/test_prepared_local_activity_client.py | 163 +++++++++++++++++++ 2 files changed, 253 insertions(+), 2 deletions(-) create mode 100644 tests/test_prepared_local_activity_client.py diff --git a/src/durable_workflow/client.py b/src/durable_workflow/client.py index c22cc0c..1afd6ff 100644 --- a/src/durable_workflow/client.py +++ b/src/durable_workflow/client.py @@ -108,10 +108,32 @@ ) _RUNTIME_EXTERNAL_PAYLOAD_ERROR_BODY_LIMIT = 64 * 1024 _PAYLOAD_COMPLETION_SCHEMA = "durable-workflow.v2.payload-completion-context.v1" +_PREPARED_PAYLOAD_COMPLETION_SCHEMA = "durable-workflow.v2.payload-completion-context.v2" _PAYLOAD_COMPLETION_HEADER = "X-Durable-Workflow-Payload-Completion" -def _payload_completion_context(path: str, body: Any) -> dict[str, Any] | None: +def _payload_completion_context(path: str, body: Any, *, allow_prepared: bool = False) -> dict[str, Any] | None: + prepared = re.fullmatch( + r"/worker/workflow-tasks/([^/]+)/local-activities/(?:(checkpoint|checkpoint-group|prepare|recover)|([^/]+)/outcome)", + path.split("?")[0], + ) + if allow_prepared and prepared is not None and isinstance(body, dict): + task_id, operation, activity_attempt_id = prepared.groups() + operation = operation or "outcome" + identity = ({"checkpoint_id": body.get("checkpoint_id")} if operation in {"checkpoint", "checkpoint-group"} + else {"activity_attempt_id": unquote(activity_attempt_id)} if activity_attempt_id is not None + else {"sequence": body.get("sequence")}) + value = next(iter(identity.values())) + owner, attempt = body.get("lease_owner"), body.get("workflow_task_attempt") + identity_valid = (type(value) is int and value > 0) if operation in {"prepare", "recover"} \ + else isinstance(value, str) and bool(value.strip()) + if (not isinstance(owner, str) or not owner.strip() or type(attempt) is not int or attempt < 1 + or not identity_valid): + return None + return {"schema": _PREPARED_PAYLOAD_COMPLETION_SCHEMA, "kind": "workflow", "task_id": unquote(task_id), + "attempt": attempt, "lease_owner": owner, + "operation": "local_activity_group_checkpoint" if operation == "checkpoint-group" + else "local_activity_" + operation, **identity} match = re.fullmatch(r"/worker/(activity|workflow|query)-tasks/([^/]+)/(complete|fail)", path.split("?")[0]) if match is None or not isinstance(body, dict): return None @@ -311,6 +333,7 @@ class _RuntimeExternalPayloadTransport: request_timeout_seconds: float status: str completion_context: bool = False + prepared_completion_context: bool = False @dataclass @@ -1796,7 +1819,8 @@ async def _request( transport=transport, uploaded={}, completion=( - _payload_completion_context(path, json) if worker and transport.completion_context else None + _payload_completion_context(path, json, allow_prepared=transport.prepared_completion_context) + if worker and transport.completion_context else None ), ) @@ -1977,6 +2001,11 @@ async def _runtime_external_payload_transport( completion_context=isinstance(completion, dict) and completion.get("schema") == _PAYLOAD_COMPLETION_SCHEMA and completion.get("header") == _PAYLOAD_COMPLETION_HEADER, + prepared_completion_context=isinstance(completion, dict) + and completion.get("prepared_schema") == _PREPARED_PAYLOAD_COMPLETION_SCHEMA + and _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )), ) self._runtime_external_payload_transport_cache = transport self._runtime_external_payload_transport_resolved = True @@ -5641,6 +5670,65 @@ async def acknowledge_activity_cancellation( raise ServerError(200, {"reason": "invalid_activity_cancellation_acknowledgement"}) return result + async def prepared_local_activity_operation( + self, + *, + task_id: str, + lease_owner: str, + workflow_task_attempt: int, + operation: str, + body: Mapping[str, Any] | None = None, + activity_attempt_id: str | None = None, + timeout_seconds: float = 5.0, + ) -> dict[str, Any]: + """Perform one prepared-local operation on its original workflow claim. + + Source protocol 1.20 only. This transport does not grant permission to + invoke a callback. The worker must discover the installed bridge and + validate its original admission receipt, deadlines and canonical history. + Every retry and payload transfer shares this total authority budget. + """ + if not _supports_cooperative_cancellation_protocol(_protocol_version_from_env( + "DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION, + )): + raise ValueError("prepared local operations require explicit worker protocol 1.20") + admission = {"checkpoint", "checkpoint-group", "prepare", "recover"} + attempts = {"control", "heartbeat", "outcome", "acknowledge-cancellation"} + if operation not in admission | attempts: + raise ValueError("unsupported prepared local operation") + identifiers = [task_id, lease_owner] + if operation in attempts: + if activity_attempt_id is None: + raise ValueError("prepared local attempt operation requires its original backend attempt identity") + identifiers.append(activity_attempt_id) + elif activity_attempt_id is not None: + raise ValueError("prepared local admission cannot carry an activity attempt identity") + if any(not isinstance(value, str) or not value.strip() or len(value.encode("utf-8")) > 255 + for value in identifiers): + raise ValueError("prepared local operations require bounded nonempty original claim identities") + if type(workflow_task_attempt) is not int or workflow_task_attempt < 1: + raise ValueError("prepared local operations require a positive original workflow task epoch") + if (isinstance(timeout_seconds, bool) or not isinstance(timeout_seconds, int | float) + or not math.isfinite(timeout_seconds) or not 0 < timeout_seconds <= 5): + raise ValueError("prepared local authority budget must be positive, finite and at most five seconds") + if body is not None and not isinstance(body, Mapping): + raise TypeError("prepared local operation body must be an object") + payload = dict(body or {}) + if "lease_owner" in payload or "workflow_task_attempt" in payload: + raise ValueError("prepared local operation body cannot replace original workflow claim authority") + payload = {"lease_owner": lease_owner, "workflow_task_attempt": workflow_task_attempt, **payload} + json_module.dumps(payload, allow_nan=False) + path = f"/worker/workflow-tasks/{quote(task_id, safe='._:-')}/local-activities/" + if activity_attempt_id is not None: + path += quote(activity_attempt_id, safe="._:-") + "/" + result = await asyncio.wait_for( + self._request("POST", path + operation, worker=True, json=payload, timeout=timeout_seconds), + timeout=timeout_seconds, + ) + if not isinstance(result, dict): + raise ServerError(200, {"reason": "invalid_prepared_local_activity_receipt"}) + return result + async def heartbeat_activity_task( self, *, diff --git a/tests/test_prepared_local_activity_client.py b/tests/test_prepared_local_activity_client.py new file mode 100644 index 0000000..b499da2 --- /dev/null +++ b/tests/test_prepared_local_activity_client.py @@ -0,0 +1,163 @@ +from __future__ import annotations + +import asyncio +import json +import time +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from durable_workflow.client import Client, _payload_completion_context +from durable_workflow.errors import ExternalPayloadError, ServerError +from durable_workflow.retry_policy import TransportRetryPolicy + + +@pytest.fixture(autouse=True) +def candidate_protocol(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + + +@pytest.mark.parametrize("operation", ["checkpoint", "checkpoint-group", "prepare", "recover", + "control", "heartbeat", "outcome", "acknowledge-cancellation"]) +async def test_operation_keeps_original_claim_and_encoded_backend_identity(operation: str) -> None: + attempt = "backend/attempt" if operation in { + "control", "heartbeat", "outcome", "acknowledge-cancellation", + } else None + async with Client("http://server", worker_token="worker", namespace="tenant") as client: + with patch.object(client, "_request", new_callable=AsyncMock, return_value={"opaque": "receipt"}) as send: + result = await client.prepared_local_activity_operation(task_id="task/one", lease_owner="original", + workflow_task_attempt=4, operation=operation, body={"checkpoint_id": "group"}, + activity_attempt_id=attempt, timeout_seconds=0.5) + assert result == {"opaque": "receipt"} + suffix = "backend%2Fattempt/" if attempt else "" + send.assert_awaited_once_with("POST", "/worker/workflow-tasks/task%2Fone/local-activities/" + suffix + operation, + worker=True, json={"lease_owner": "original", "workflow_task_attempt": 4, "checkpoint_id": "group"}, + timeout=0.5) + + +async def test_actual_transport_uses_worker_credential_namespace_and_candidate_header() -> None: + reply = httpx.Response(200, json={"active": False}, request=httpx.Request("POST", "http://server")) + async with Client("http://server", control_token="control", worker_token="worker", namespace="tenant") as client: + with patch.object(client._http, "request", new_callable=AsyncMock, return_value=reply) as send: + await client.prepared_local_activity_operation(task_id="task/one", lease_owner="original", + workflow_task_attempt=4, operation="control", activity_attempt_id="backend/attempt") + args, kwargs = send.await_args + assert args == ("POST", "/api/worker/workflow-tasks/task%2Fone/local-activities/backend%2Fattempt/control") + assert kwargs["headers"]["Authorization"] == "Bearer worker" + assert kwargs["headers"]["X-Namespace"] == "tenant" + assert kwargs["headers"]["X-Durable-Workflow-Protocol-Version"] == "1.20" + + +@pytest.mark.parametrize("change", [ + {"task_id": ""}, {"lease_owner": " "}, {"workflow_task_attempt": True}, {"workflow_task_attempt": 0}, + {"operation": "guess"}, {"activity_attempt_id": "unexpected"}, {"body": {"lease_owner": "replacement"}}, + {"body": {"workflow_task_attempt": 5}}, {"body": {"not_finite": float("nan")}}, + {"timeout_seconds": 0}, {"timeout_seconds": float("inf")}, {"timeout_seconds": 5.1}, +]) +async def test_invalid_claim_or_authority_never_starts_transport(change: dict[str, Any]) -> None: + arguments = {"task_id": "task", "lease_owner": "original", "workflow_task_attempt": 4, + "operation": "prepare", **change} + async with Client("http://server") as client: + with patch.object(client, "_request", new_callable=AsyncMock) as send: + with pytest.raises((ValueError, TypeError)): + await client.prepared_local_activity_operation(**arguments) + send.assert_not_awaited() + + +async def test_default_protocol_refuses_before_io(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION") + async with Client("http://server") as client: + with patch.object(client, "_request", new_callable=AsyncMock) as send: + with pytest.raises(ValueError, match="1.20"): + await client.prepared_local_activity_operation(task_id="task", lease_owner="original", + workflow_task_attempt=4, operation="prepare") + send.assert_not_awaited() + + +async def test_one_total_budget_cancels_pending_payload_and_retry_work() -> None: + cancelled = False + + async def pending(*args: Any, **kwargs: Any) -> None: + nonlocal cancelled + try: + await asyncio.sleep(60) + finally: + cancelled = True + + async with Client("http://server") as client: + with patch.object(client, "_request", side_effect=pending): + started = time.monotonic() + with pytest.raises(asyncio.TimeoutError): + await client.prepared_local_activity_operation(task_id="task", lease_owner="original", + workflow_task_attempt=4, operation="prepare", timeout_seconds=0.08) + assert time.monotonic() - started < 0.3 + assert cancelled + + +@pytest.mark.parametrize("receipt", [None, [], "unproved"]) +async def test_non_object_receipt_is_explicitly_refused(receipt: Any) -> None: + async with Client("http://server") as client: + with patch.object(client, "_request", new_callable=AsyncMock, return_value=receipt): + with pytest.raises(ServerError) as failure: + await client.prepared_local_activity_operation(task_id="task", lease_owner="original", + workflow_task_attempt=4, operation="prepare") + assert failure.value.reason() == "invalid_prepared_local_activity_receipt" + + +def test_payload_admission_identity_requires_negotiation_and_original_epoch() -> None: + body = {"lease_owner": "original", "workflow_task_attempt": 4, "checkpoint_id": "atomic-group"} + route = "/worker/workflow-tasks/task%2Fone/local-activities/checkpoint-group" + assert _payload_completion_context(route, body) is None + assert _payload_completion_context(route, body, allow_prepared=True) == { + "schema": "durable-workflow.v2.payload-completion-context.v2", "kind": "workflow", + "task_id": "task/one", "attempt": 4, "lease_owner": "original", + "operation": "local_activity_group_checkpoint", "checkpoint_id": "atomic-group", + } + for invalid in [ + {**body, "lease_owner": ""}, {**body, "workflow_task_attempt": True}, {**body, "checkpoint_id": ""}, + ]: + assert _payload_completion_context(route, invalid, allow_prepared=True) is None + + +@pytest.mark.parametrize("prepared_supported,state", [(True, "draining"), (False, "draining"), (True, "fenced")]) +async def test_group_draining_uploads_bind_each_payload_slot_and_original_claim( + prepared_supported: bool, state: str, +) -> None: + from durable_workflow import serializer + from tests.test_runtime_external_payload_transport import CompletionPayloadServer, runtime_client + + class PreparedPayloadServer(CompletionPayloadServer): + def cluster_info(self) -> dict[str, Any]: + info = super().cluster_info() + context = info["namespace"]["external_payload_storage"]["transport"]["upload"]["completion_context"] + if prepared_supported: + context["prepared_schema"] = "durable-workflow.v2.payload-completion-context.v2" + return info + + server = PreparedPayloadServer(state=state) + async with runtime_client(server, retry_policy=TransportRetryPolicy(max_attempts=1)) as client: + call = client.prepared_local_activity_operation(task_id="task/one", lease_owner="original", + workflow_task_attempt=4, operation="checkpoint-group", body={"checkpoint_id": "group", "commands": [ + {"type": "start_child_workflow", "input": serializer.envelope("child" * 100)}, + {"type": "prepare_local_activity", "arguments": serializer.envelope("local" * 100)}, + ]}) + if not prepared_supported or state == "fenced": + with pytest.raises((ServerError, ExternalPayloadError)): + await call + assert len(server.upload_requests) == 1 + assert "X-Durable-Workflow-Payload-Completion" not in server.upload_requests[0].headers + assert server.requests == [] + return + await call + assert len(server.upload_requests) == 4 + for index, slot in [(0, ["commands", 0, "input"]), (2, ["commands", 1, "arguments"])]: + first, bound = server.upload_requests[index:index + 2] + assert first.content == bound.content + assert bound.headers["authorization"] == first.headers["authorization"] + assert json.loads(bound.headers["X-Durable-Workflow-Payload-Completion"]) == { + "schema": "durable-workflow.v2.payload-completion-context.v2", "kind": "workflow", + "task_id": "task/one", "attempt": 4, "lease_owner": "original", + "operation": "local_activity_group_checkpoint", "checkpoint_id": "group", "slot": slot, + } From 5150044c3bec86b8d83e206f5468dbe8a0057d8e Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 11:26:36 +0000 Subject: [PATCH 24/41] feat: admit and supervise prepared local callbacks --- .github/workflows/ci.yml | 2 +- docs/cooperative-cancellation-design.md | 37 +- .../_prepared_local_activity.py | 340 ++++++++++++ src/durable_workflow/worker.py | 220 +++++++- src/durable_workflow/workflow.py | 88 +++- .../prepared-local-cold-results.json | 23 + tests/integration/prepared_worker.py | 26 + .../test_prepared_local_activity.py | 215 ++++++++ tests/test_prepared_local_activity_worker.py | 483 ++++++++++++++++++ tests/test_replay_regression_corpus.py | 9 + 10 files changed, 1429 insertions(+), 14 deletions(-) create mode 100644 src/durable_workflow/_prepared_local_activity.py create mode 100644 tests/fixtures/replay_regressions/prepared-local-cold-results.json create mode 100644 tests/integration/prepared_worker.py create mode 100644 tests/integration/test_prepared_local_activity.py create mode 100644 tests/test_prepared_local_activity_worker.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 1182035..134f726 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -238,7 +238,7 @@ jobs: sdk-python/integration-package-provenance.txt sdk-python/integration-images.json if-no-files-found: warn - retention-days: 7 + retention-days: ${{ inputs.cooperative_qualification && 90 || 7 }} - name: Emit integration diagnostics if: failure() working-directory: sdk-python diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index bbd1fa9..832a10e 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -108,11 +108,44 @@ application entry point with `if __name__ == "__main__"`. Captured memory change are local to the callback process. Open process-local connections in the callback. Registration and remote polling refuse incompatible handler or interceptor definitions before claiming work, naming the worker and activity. -Cooperative local callbacks still require their separate supervision and receipt -model. Legacy worker protocol 1.19 continues using its existing execution. +Prepared local callbacks use the same physical process ownership with separate +durable admission and receipts. Legacy worker protocol 1.19 continues using its existing execution. Cooperating downstream systems still need idempotency or reconciliation for effects already performed. +## Prepared sequential local callbacks + +The source candidate can explicitly request both `cooperative_cancellation` and +`prepared_local_activities` in Worker capabilities. Registration requires source +protocol 1.20 and actual Server discovery of its installed admission bridge. +The manifest advertises `durable_sequential_admission`. Python refuses prepared +local parallel and selection groups until its atomic group consumer exists. + +Replay captures the authored local call and sequence before application code +runs. Earlier side effects, version markers and metadata commands obtain a +retained-claim checkpoint, followed by canonical history refresh. The Server +then creates the local execution and original attempt. The worker validates its +workflow claim, epoch, owner, backend IDs, nonce, fixed deadlines and cleanup +authority before spawning. Native owns retries, backoff and execution timeouts. + +Independent control renews ownership without creating application heartbeat +history or advancing its timeout. Only a real callback heartbeat may do those +things. Requests, retries, payload transfer and result encoding share the +original conservative authority budget. A cancellation fence physically stops +and joins callback and supervisor before the stop receipt. An unconfirmed stop +retains workflow capacity and prevents successful worker deregistration. + +Canonical outcome history supplies the next replay value. A lost or malformed +receipt abandons the claim. Cold replay skips completed callbacks and requests +Native recovery of unfinished Started attempts. Recovery records unknown stop +and may release the claim for a durable retry, without claiming that the +replacement observed the original callback stop. + +Cleanup local calls require a shield after canonical delivery. Admission and +control preserve its original local request ID, root ID, delivery history event +ID and deadline. No replacement or duplicate request grants a fresh cleanup +budget. Connected exact-source qualification remains separate from publication. + ## Remaining qualification Rust policy parity, portable activity policies, nested scopes and deterministic diff --git a/src/durable_workflow/_prepared_local_activity.py b/src/durable_workflow/_prepared_local_activity.py new file mode 100644 index 0000000..1e106f0 --- /dev/null +++ b/src/durable_workflow/_prepared_local_activity.py @@ -0,0 +1,340 @@ +"""Source-only prepared local admission and supervised callback execution. + +Native owns retry policy, execution deadlines and terminal history. Control +observes authority independently of application heartbeats. A lost receipt +abandons the claim and never authorizes another callback or publication. +""" + +from __future__ import annotations + +import asyncio +import re +import time +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import Any + +from . import serializer +from ._activity_process import CallbackInvocation, SupervisedCallback +from .cancellation import CancellationContext +from .client import Client +from .workflow import LocalActivityExecutionAborted + +_FIXED_DEADLINES = ("start_to_close_deadline_at", "schedule_to_close_deadline_at") + + +def _require(condition: bool, message: str) -> None: + if not condition: + raise LocalActivityExecutionAborted(message) + + +def _text(value: Any) -> str: + _require(isinstance(value, str) and bool(value.strip()), "prepared receipt lacks a nonempty identity") + return str(value) + + +def _timestamp(value: Any) -> datetime: + _require(isinstance(value, str) and re.fullmatch( + r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:\d{2})", value, + ) is not None, "prepared receipt requires an ISO timestamp with timezone") + try: + return datetime.fromisoformat(str(value).replace("Z", "+00:00")) + except ValueError as error: + raise LocalActivityExecutionAborted("prepared receipt timestamp is invalid") from error + + +def _same_deadline(actual: Any, expected: Any) -> bool: + return actual is None if expected is None else _timestamp(actual) == _timestamp(expected) + + +class PreparedCancellationObserved(LocalActivityExecutionAborted): + """Cancellation was fenced and its original callback physically joined.""" + + +@dataclass +class PreparedAttempt: + task_id: str + run_id: str + owner: str + epoch: int + nonce: str + execution_id: str + attempt_id: str + attempt_number: int + deadlines: dict[str, Any] + heartbeat_timeout: int | None + cleanup: Mapping[str, str] | None + clock_origin: float + authority_deadline: float + + @classmethod + def admitted( + cls, receipt: Mapping[str, Any], *, task_id: str, run_id: str, owner: str, epoch: int, nonce: str, + heartbeat_timeout: int | None, cleanup: Mapping[str, str] | None, request_started: float, + ) -> PreparedAttempt: + _require( + receipt.get("prepared") is True and type(receipt.get("duplicate")) is bool + and "reason" in receipt and receipt["reason"] is None + and receipt.get("workflow_task_id") == task_id and receipt.get("lease_owner") == owner + and type(receipt.get("workflow_task_attempt")) is int and receipt["workflow_task_attempt"] == epoch + and receipt.get("worker_attempt_id") == nonce + and type(receipt.get("attempt_number")) is int and receipt["attempt_number"] > 0, + "local admission changed its original workflow claim or nonce", + ) + deadlines = {field: receipt.get(field) for field in (*_FIXED_DEADLINES, "heartbeat_deadline_at")} + attempt = cls(task_id, run_id, owner, epoch, nonce, _text(receipt.get("activity_execution_id")), + _text(receipt.get("activity_attempt_id")), receipt["attempt_number"], deadlines, + heartbeat_timeout, cleanup, request_started - _timestamp(receipt.get("server_time")).timestamp(), + request_started) + attempt.validate_cleanup(receipt) + server_time = _timestamp(receipt.get("server_time")) + _require(_timestamp(receipt.get("lease_expires_at")) > server_time, "local admission has an expired lease") + cleanup_deadline = _timestamp(cleanup["cleanup_deadline_at"]) if cleanup is not None else None + for field, deadline in deadlines.items(): + _require(field in receipt and (deadline is None or _timestamp(deadline) > server_time), + "local admission omitted or exhausted an execution deadline") + if deadline is not None and cleanup_deadline is not None: + _require(_timestamp(deadline) <= cleanup_deadline, + "local admission extended its original cleanup budget") + heartbeat = deadlines["heartbeat_deadline_at"] + if heartbeat_timeout is not None: + _require(heartbeat is not None + and _timestamp(heartbeat) <= server_time + timedelta(seconds=heartbeat_timeout), + "local admission changed its authored application heartbeat timeout") + else: + _require(_same_deadline(heartbeat, cleanup_deadline.isoformat() if cleanup_deadline is not None else None), + "local admission invented an application heartbeat deadline") + if cleanup_deadline is not None: + _require(cleanup_deadline > server_time, "local admission exhausted its original cleanup budget") + attempt.accept_budget(receipt, request_started) + attempt.remaining() + return attempt + + def validate_cleanup(self, receipt: Mapping[str, Any]) -> None: + actual = receipt.get("cancellation_cleanup") + if self.cleanup is None: + _require(actual is None, "local receipt invented cleanup authority") + return + fields = {"request_id", "root_request_id", "delivery_history_event_id", "cleanup_deadline_at"} + if not isinstance(actual, Mapping) or set(actual) != fields or set(self.cleanup) != fields: + raise LocalActivityExecutionAborted("local receipt changed canonical cleanup authority") + for field in fields: + _require(_same_deadline(actual[field], self.cleanup[field]) if field == "cleanup_deadline_at" + else _text(actual[field]) == _text(self.cleanup[field]), + "local receipt changed canonical cleanup authority") + + def validate_identity(self, receipt: Mapping[str, Any]) -> None: + _require( + receipt.get("workflow_task_id") == self.task_id + and type(receipt.get("workflow_task_attempt")) is int and receipt["workflow_task_attempt"] == self.epoch + and receipt.get("activity_execution_id") == self.execution_id + and receipt.get("activity_attempt_id") == self.attempt_id + and ("lease_owner" not in receipt or receipt["lease_owner"] == self.owner), + "prepared receipt belongs to a different attempt or workflow claim", + ) + + def validate_control(self, receipt: Mapping[str, Any], *, heartbeat: bool = False) -> None: + self.validate_identity(receipt) + self.validate_cleanup(receipt) + active = receipt.get("active") + _require( + type(active) is bool and type(receipt.get("renewed")) is bool + and receipt.get("lease_owner") == self.owner + and type(receipt.get("stop_required")) is bool and receipt["stop_required"] is not active + and "reason" in receipt + and (receipt["reason"] is None and receipt["renewed"] is (not heartbeat) if active + else isinstance(receipt["reason"], str) and bool(receipt["reason"]) and receipt["renewed"] is False), + "local control lacks live authority or an explicit stop", + ) + recorded = receipt.get("heartbeat_recorded") + _require( + recorded is active and (isinstance(receipt.get("heartbeat_history_event_id"), str) + and bool(receipt["heartbeat_history_event_id"].strip()) if active + else receipt.get("heartbeat_history_event_id") is None) if heartbeat + else recorded is False and receipt.get("heartbeat_history_event_id") is None, + "supervisor control and application heartbeat receipts must remain distinct", + ) + server_time = _timestamp(receipt.get("server_time")) + for field in _FIXED_DEADLINES: + _require(field in receipt and _same_deadline(receipt[field], self.deadlines[field]), + "local control changed an original execution deadline") + _require("heartbeat_deadline_at" in receipt, "local control omitted the application heartbeat deadline") + updated = receipt["heartbeat_deadline_at"] + if heartbeat and active and self.heartbeat_timeout is not None: + _require(_timestamp(self.deadlines["heartbeat_deadline_at"]) > server_time + and _timestamp(updated) >= _timestamp(self.deadlines["heartbeat_deadline_at"]) + and _timestamp(updated) <= server_time + timedelta(seconds=self.heartbeat_timeout) + and (self.cleanup is None + or _timestamp(updated) <= _timestamp(self.cleanup["cleanup_deadline_at"])), + "application heartbeat extended a fixed budget or revived an expired deadline") + else: + _require(_same_deadline(updated, self.deadlines["heartbeat_deadline_at"]), + "local control changed an unacknowledged application heartbeat deadline") + if active: + for field in ("lease_expires_at", "workflow_lease_expires_at"): + _require(_timestamp(receipt.get(field)) > server_time, "local control returned an expired lease") + for deadline in (receipt[field] for field in (*_FIXED_DEADLINES, "heartbeat_deadline_at")): + _require(deadline is None or _timestamp(deadline) > server_time, + "local control returned active after an execution deadline") + if self.cleanup is not None: + _require(_timestamp(self.cleanup["cleanup_deadline_at"]) > server_time, + "local control returned active after the original cleanup deadline") + if heartbeat and active: + self.deadlines["heartbeat_deadline_at"] = updated + + def validate_outcome(self, receipt: Mapping[str, Any]) -> None: + self.validate_identity(receipt) + retry = bool(receipt.get("event_type") == "ActivityRetryScheduled") + created = receipt.get("created_task_ids") + _require( + receipt.get("recorded") is True and type(receipt.get("duplicate")) is bool + and "reason" in receipt and receipt["reason"] is None + and receipt.get("workflow_run_id") == self.run_id and receipt.get("worker_attempt_id") == self.nonce + and receipt.get("event_type") in {"ActivityCompleted", "ActivityFailed", "ActivityTimedOut", + "ActivityRetryScheduled"} + and receipt.get("claim_released") is retry and isinstance(created, list) and len(created) == int(retry), + "local outcome lacks a canonical receipt for its original attempt", + ) + _text(receipt.get("event_id")) + _timestamp(receipt.get("recorded_at")) + for identity in receipt["created_task_ids"]: + _text(identity) + + def accept_budget(self, receipt: Mapping[str, Any], started: float) -> None: + server_time = _timestamp(receipt.get("server_time")).timestamp() + values = [receipt.get(field) for field in ( + "lease_expires_at", "workflow_lease_expires_at", *_FIXED_DEADLINES, "heartbeat_deadline_at", + )] + if self.cleanup is not None: + values.append(self.cleanup["cleanup_deadline_at"]) + self.authority_deadline = min( + min(self.clock_origin + _timestamp(value).timestamp(), + started + _timestamp(value).timestamp() - server_time) + for value in values if value is not None + ) + + def remaining(self) -> float: + remaining = self.authority_deadline - time.monotonic() + _require(remaining > 0, "prepared local original authority budget expired") + return min(5.0, remaining) + + +class PreparedLocalRunner: + def __init__( + self, client: Client, attempt: PreparedAttempt, *, observe: Callable[[Mapping[str, Any]], Any], + shutdown: asyncio.Event, + ) -> None: + self.client = client + self.attempt = attempt + self.observe = observe + self.shutdown = shutdown + self.stop_request: CancellationContext | None = None + self.lock = asyncio.Lock() + self.callback: SupervisedCallback | None = None + + async def operation(self, operation: str, body: Mapping[str, Any]) -> dict[str, Any]: + return await self.client.prepared_local_activity_operation( + task_id=self.attempt.task_id, lease_owner=self.attempt.owner, workflow_task_attempt=self.attempt.epoch, + activity_attempt_id=self.attempt.attempt_id, operation=operation, body=body, + timeout_seconds=self.attempt.remaining(), + ) + + async def control(self, details: dict[str, Any] | None = None, *, heartbeat: bool = False) -> None: + async with self.lock: + _require(not self.shutdown.is_set(), "worker shutdown abandoned its prepared local claim") + started = time.monotonic() + receipt = await self.operation("heartbeat" if heartbeat else "control", + {"details": details} if heartbeat else {"renew_lease": True}) + self.attempt.validate_control(receipt, heartbeat=heartbeat) + if receipt["active"]: + self.attempt.accept_budget(receipt, started) + self.attempt.remaining() + return + if receipt["reason"] not in {"cancellation_requested", "cancellation_deadline_expired"}: + raise LocalActivityExecutionAborted("prepared local callback lost its original authority") + try: + context = CancellationContext.from_dict(receipt.get("cancellation_request", {})) + except (ValueError, TypeError, AttributeError) as error: + raise LocalActivityExecutionAborted("prepared stop lacks canonical cancellation context") from error + _require(context.lineage[-1].workflow_run_id == self.attempt.run_id and receipt.get("fenced") is True, + "prepared cancellation did not fence this original callback") + _text(receipt.get("cancellation_history_event_id")) + self.observe({**context.to_dict(), "history_refresh_page_token": receipt.get("history_refresh_page_token")}) + self.stop_request = context + raise PreparedCancellationObserved("prepared callback observed cooperative cancellation") + + async def execute(self, invocation: CallbackInvocation, processes: set[SupervisedCallback]) -> dict[str, Any]: + async def observe() -> None: + while True: + await asyncio.sleep(min(1.0, self.attempt.remaining())) + await self.control() + + async def heartbeat(details: dict[str, Any] | None) -> None: + await self.control(details, heartbeat=True) + + try: + await self.control() + except PreparedCancellationObserved: + # Admission was fenced before any process existed. + await self.acknowledge_stop() + raise + callback = SupervisedCallback(invocation) + self.callback = callback + processes.add(callback) + background: list[asyncio.Task[Any]] = [] + try: + await asyncio.wait_for(callback.start(), timeout=self.attempt.remaining()) + result = asyncio.create_task(callback.result(heartbeat)) + observer = asyncio.create_task(observe()) + shutdown = asyncio.create_task(self.shutdown.wait()) + background = [result, observer, shutdown] + done, _ = await asyncio.wait(background, return_when=asyncio.FIRST_COMPLETED) + if shutdown in done: + raise LocalActivityExecutionAborted("worker shutdown abandoned its prepared local claim") + if observer in done: + await observer + raise LocalActivityExecutionAborted("prepared local authority observer stopped") + outcome = await result + await self.control() + report: dict[str, Any] + if outcome.failure is None: + report = {"outcome": "completed", "result": serializer.envelope(outcome.value), + "payload_codec": serializer.AVRO_CODEC} + else: + failure = outcome.failure + report = {"outcome": "failed", "message": failure.message, "exception_type": failure.failure_type, + "non_retryable": failure.non_retryable or failure.cancelled} + # Cancel the observer before publication. No control may race a + # committed result and mistake its closed attempt for lost authority. + observer.cancel() + await asyncio.gather(observer, return_exceptions=True) + async with self.lock: + receipt = await self.operation("outcome", {"report": report}) + self.attempt.validate_outcome(receipt) + return receipt + finally: + for pending in background: + if not pending.done(): + pending.cancel() + await asyncio.gather(*background, return_exceptions=True) + if not callback.stopped: + await callback.stop() + _require(callback.stopped, "prepared callback stop remains unconfirmed") + processes.discard(callback) + await self.acknowledge_stop() + + async def acknowledge_stop(self) -> None: + if self.stop_request is None: + return + # A physical join, or the absence of any spawned callback, is proved + # before this diagnostic receipt. It grants no execution authority. + receipt = await self.client.prepared_local_activity_operation( + task_id=self.attempt.task_id, lease_owner=self.attempt.owner, + workflow_task_attempt=self.attempt.epoch, activity_attempt_id=self.attempt.attempt_id, + operation="acknowledge-cancellation", body={"request_id": self.stop_request.request_id}, + ) + _require(receipt.get("acknowledged") is True and type(receipt.get("duplicate")) is bool + and "reason" in receipt and receipt["reason"] is None, + "joined prepared callback stop was not acknowledged") + _text(receipt.get("history_event_id")) diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 86738b6..329e193 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -39,6 +39,7 @@ from . import serializer from ._activity_process import CallbackFailure, CallbackInvocation, CallbackProcessLost, SupervisedCallback from ._cooperative_cancellation import CancellationRequest, read_cancellation_history +from ._prepared_local_activity import PreparedAttempt, PreparedCancellationObserved, PreparedLocalRunner from .activity import ActivityContext, ActivityInfo, _set_context from .auth_composition import ( AUTH_COMPOSITION_CONTRACT_SCHEMA, @@ -46,6 +47,7 @@ AuthCompositionContractError, parse_auth_composition_contract, ) +from .cancellation import CancellationContext from .client import ( CONTROL_PLANE_REQUEST_CONTRACT_SCHEMA, CONTROL_PLANE_REQUEST_CONTRACT_VERSION, @@ -1049,6 +1051,11 @@ def __init__( self.activities = {_activity_name(a): a for a in activities} self.capabilities = tuple(dict.fromkeys(capability.strip() for capability in capabilities)) self._cooperative_cancellation_supported = False + self._prepared_local_activities_supported = False + if "prepared_local_activity_groups" in self.capabilities: + raise ValueError("Python has no prepared local group consumer yet") + if "prepared_local_activities" in self.capabilities and "cooperative_cancellation" not in self.capabilities: + raise ValueError("prepared local activities require cooperative_cancellation capability") if any(not capability for capability in self.capabilities): raise ValueError("worker capabilities must be non-empty strings") self.worker_id = worker_id or f"py-worker-{uuid.uuid4().hex[:8]}" @@ -1080,6 +1087,8 @@ def __init__( self._remote_activity_executor: ThreadPoolExecutor | None = None self._remote_activity_threads: set[Future[Any]] = set() self._remote_activity_processes: set[SupervisedCallback] = set() + self._prepared_local_activity_processes: set[SupervisedCallback] = set() + self._abandoned_prepared_local_claims: set[str] = set() self._remote_activity_thread_slots = asyncio.Semaphore(max_concurrent_activity_tasks) self._wf_semaphore = asyncio.Semaphore(max_concurrent_workflow_tasks) self._act_semaphore = asyncio.Semaphore(max_concurrent_activity_tasks) @@ -1225,6 +1234,15 @@ async def _register(self) -> None: ) if "cooperative_cancellation" in self.capabilities and not self._cooperative_cancellation_supported: raise RuntimeError("cooperative cancellation requires explicit compatible runtime and worker protocol 1.20") + self._prepared_local_activities_supported = ( + "prepared_local_activities" in self.capabilities and self._cooperative_cancellation_supported + and isinstance(server_capabilities, Mapping) + and server_capabilities.get("prepared_local_activities") is True + ) + if "prepared_local_activities" in self.capabilities and not self._prepared_local_activities_supported: + raise RuntimeError( + "prepared_local_activity_not_supported: Server must advertise its installed admission bridge", + ) self._validate_cooperative_activity_handlers() self._query_tasks_supported = _server_supports_query_tasks(info) self._workflow_memo_updates_supported = _server_supports_workflow_memo_updates(info) @@ -1274,7 +1292,13 @@ async def _register(self) -> None: max_concurrent_worker_sessions=self.max_concurrent_worker_sessions, build_id=self.build_id, capabilities=capabilities, - capability_manifest=PORTABLE_WORKER_AFFINITY_CAPABILITY_MANIFEST, + capability_manifest={ + **PORTABLE_WORKER_AFFINITY_CAPABILITY_MANIFEST, + **({"prepared_local_activities": { + "supported": True, "minimum_protocol_version": "1.20", + "implementation": "durable_sequential_admission", + }} if self._prepared_local_activities_supported else {}), + }, task_slots=self._current_task_slots(), process_metrics=self._current_process_metrics(), ) @@ -1506,7 +1530,8 @@ async def _replay_workflow_claim( if observation is not None: observed = self._observe_workflow_cancellation(task, observation) history = await self._refresh_cancellation_history(task, observed) - for _ in range(3): + delivery_attempts = 0 + for _ in range(1000): state = read_cancellation_history( history, run_id=task.get("run_id", ""), observation=task.get("cancellation_request"), ) @@ -1525,7 +1550,34 @@ async def _replay_workflow_claim( external_storage_cache=self.external_storage_cache, cancel_requested=bool(task.get("cancel_requested", False)) and state.request is None, cancellation_request=task.get("cancellation_request"), local_activity_executor=execute_local, + prepare_local_activities=self._prepared_local_activities_supported, ) + if outcome.prepared_local_activity is not None: + history = await self._execute_prepared_local_activity(task, history, outcome) + continue + except PreparedCancellationObserved as error: + # Control carries the root request time. Workflow observations + # carry this run's local admission time, which can be later in a + # cascade. Read the actual task observation instead of replacing + # either timestamp with the other. + expected = CancellationContext.from_dict(task["_prepared_cancellation_context"]) + with contextlib.suppress(_CooperativeCancellationObserved): + await self._renew_local_workflow_lease(task) + observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) + if observed.request_id != expected.request_id or datetime.fromisoformat( + observed.cleanup_deadline_at.replace("Z", "+00:00"), + ) != expected.deadline: + raise LocalActivityExecutionAborted( + "prepared stop changed the original request or deadline", + ) from error + history = await self._refresh_cancellation_history(task, observed) + committed = read_cancellation_history(history, run_id=task.get("run_id", ""), + observation=task.get("cancellation_request")) + if committed.request is None or committed.request.context != expected: + raise LocalActivityExecutionAborted( + "prepared stop context differs from canonical request history", + ) from error + continue except _CooperativeCancellationObserved: observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) history = await self._refresh_cancellation_history(task, observed) @@ -1535,6 +1587,9 @@ async def _replay_workflow_claim( # Earlier authored commands must commit on this claim first. # Completion releases the claim; a successor replays their durable results. return outcome, history + delivery_attempts += 1 + if delivery_attempts > 3: + raise LocalActivityExecutionAborted("cancellation replay did not converge on its canonical delivery") observed = self._observe_workflow_cancellation(task, task.get("cancellation_request")) delivery_error: Exception | None = None try: @@ -1559,7 +1614,152 @@ async def _replay_workflow_claim( raise LocalActivityExecutionAborted( "delivery was not proved by matching canonical history", ) from delivery_error - raise LocalActivityExecutionAborted("workflow cancellation replay did not converge on its canonical delivery") + raise LocalActivityExecutionAborted("workflow exceeded the prepared local replay admission limit") + + async def _refresh_prepared_local_history( + self, task: dict[str, Any], receipt: Mapping[str, Any], + ) -> list[dict[str, Any]]: + token = receipt.get("history_refresh_page_token") + if not isinstance(token, str) or not token.strip(): + raise LocalActivityExecutionAborted("prepared operation lacks a Server-issued canonical history cursor") + return await self._load_workflow_claim_history(task, first_page_token=token) + + async def _execute_prepared_local_activity( + self, task: dict[str, Any], history: list[dict[str, Any]], outcome: ReplayOutcome, + ) -> list[dict[str, Any]]: + call = outcome.prepared_local_activity + if call is None: + raise LocalActivityExecutionAborted("prepared replay did not identify its authored local call") + task_id = task["task_id"] + epoch = task.get("workflow_task_attempt", 1) + codec = _validate_payload_codec(task.get("payload_codec")) or serializer.AVRO_CODEC + try: + if outcome.commands: + commands = commands_to_server_commands(outcome.commands, self.task_queue, payload_codec=codec) + start = call.sequence - len(commands) + if start < 1 or any(command["type"] not in { + "record_side_effect", "record_version_marker", "upsert_memo", "upsert_search_attributes", + } for command in commands): + raise LocalActivityExecutionAborted("prepared prefix has no supported authored sequence range") + if any(command["type"] == "upsert_memo" for command in commands) and ( + not self._workflow_memo_updates_supported + ): + raise LocalActivityExecutionAborted( + "prepared memo prefix requires negotiated workflow memo updates", + ) + checkpoint = hashlib.sha256(json.dumps( + [task_id, self.worker_id, epoch, start, commands], sort_keys=True, allow_nan=False, + ).encode()).hexdigest() + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, operation="checkpoint", + body={"checkpoint_id": checkpoint, "start_sequence": start, "commands": commands}, + ) + if (receipt.get("checkpointed") is not True or type(receipt.get("duplicate")) is not bool + or receipt.get("checkpoint_id") != checkpoint + or receipt.get("task_id") != task_id or receipt.get("workflow_run_id") != task["run_id"] + or type(receipt.get("workflow_task_attempt")) is not int + or receipt["workflow_task_attempt"] != epoch or receipt.get("lease_owner") != self.worker_id + or receipt.get("start_sequence") != start or receipt.get("next_sequence") != call.sequence + or "reason" not in receipt + or receipt["reason"] is not None): + raise LocalActivityExecutionAborted("prepared prefix lacks an original-claim checkpoint receipt") + return await self._refresh_prepared_local_history(task, receipt) + descriptor = call.descriptor(codec) + if call.recover: + original = next((event.get("payload", {}) for event in reversed(history) + if event.get("event_type", event.get("type")) == "ActivityStarted" + and event.get("payload", {}).get("sequence", event.get("payload", {}).get( + "workflow_sequence")) == call.sequence), None) + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, + operation="recover", body={"sequence": call.sequence, "descriptor": descriptor}, + ) + kind = receipt.get("event_type") + retry = bool(kind == "ActivityRetryScheduled") + created = receipt.get("created_task_ids") + if (receipt.get("recovered") is not True or type(receipt.get("duplicate")) is not bool + or "reason" not in receipt or receipt["reason"] is not None + or receipt.get("workflow_task_id") != task_id or original is None + or not isinstance(receipt.get("activity_execution_id"), str) + or not receipt["activity_execution_id"].strip() + or not isinstance(receipt.get("activity_attempt_id"), str) + or not receipt["activity_attempt_id"].strip() + or receipt["activity_execution_id"] != original.get("activity_execution_id") + or receipt["activity_attempt_id"] != original.get("activity_attempt_id") + or receipt.get("callback_stop_state") != "unknown" + or kind not in {"ActivityRetryScheduled", "ActivityFailed", "ActivityTimedOut"} + or receipt.get("claim_released") is not retry + or not isinstance(created, list) or len(created) != int(retry) + or any(not isinstance(value, str) or not value.strip() for value in created) + or not isinstance(receipt.get("event_id"), str) or not receipt["event_id"].strip()): + raise LocalActivityExecutionAborted("prepared recovery lacks its canonical unknown-stop receipt") + if retry: + raise _WorkflowClaimDeferred("durable prepared local retry released the original workflow claim") + refreshed = await self._refresh_prepared_local_history(task, receipt) + event = next((event for event in refreshed if event.get("id") == receipt["event_id"]), None) + payload = event.get("payload", {}) if event is not None else {} + recovery = payload.get("local_recovery", {}) + if (event is None or event.get("event_type", event.get("type")) != kind + or payload.get("sequence", payload.get("workflow_sequence")) != call.sequence + or payload.get("activity_execution_id") != receipt["activity_execution_id"] + or payload.get("activity_attempt_id") != receipt["activity_attempt_id"] + or recovery.get("workflow_task_id") != task_id or recovery.get("workflow_task_attempt") != epoch + or recovery.get("lease_owner") != self.worker_id + or recovery.get("callback_stop_state") != "unknown"): + raise LocalActivityExecutionAborted( + "prepared recovery is absent from this claim's canonical history", + ) + return refreshed + nonce = hashlib.sha256(json.dumps( + [task_id, task["run_id"], self.worker_id, epoch, call.sequence], + ).encode()).hexdigest() + started = time.monotonic() + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, operation="prepare", + body={"sequence": call.sequence, "worker_attempt_id": nonce, "descriptor": descriptor}, + ) + attempt = PreparedAttempt.admitted( + receipt, task_id=task_id, run_id=task["run_id"], owner=self.worker_id, epoch=epoch, nonce=nonce, + heartbeat_timeout=call.command.heartbeat_timeout, cleanup=call.cleanup, request_started=started, + ) + handler = self.activities.get(call.command.activity_type) + if handler is None: + raise LocalActivityExecutionAborted("prepared local activity has no registered handler") + runner = PreparedLocalRunner( + self.client, attempt, shutdown=self._local_activity_shutdown, + observe=lambda value: task.update({"_prepared_cancellation_context": dict(value)}), + ) + try: + receipt = await runner.execute(CallbackInvocation( + handler, tuple(call.command.arguments), ActivityInfo( + task_id, call.command.activity_type, attempt.attempt_id, attempt.attempt_number, + self.task_queue, self.worker_id, + ), {**task, "activity_attempt_id": attempt.attempt_id}, self.interceptors, + ), self._prepared_local_activity_processes) + finally: + if runner.callback is not None and runner.callback in self._prepared_local_activity_processes: + self._abandoned_prepared_local_claims.add(task_id) + if receipt["claim_released"]: + raise _WorkflowClaimDeferred("durable prepared local retry released the original workflow claim") + refreshed = await self._refresh_prepared_local_history(task, receipt) + event = next((event for event in refreshed if event.get("id") == receipt["event_id"]), None) + payload = event.get("payload", {}) if event is not None else {} + if (event is None or event.get("event_type", event.get("type")) != receipt["event_type"] + or payload.get("sequence", payload.get("workflow_sequence")) != call.sequence + or payload.get("activity_execution_id") != attempt.execution_id + or payload.get("activity_attempt_id") != attempt.attempt_id): + raise LocalActivityExecutionAborted("prepared outcome is absent from this claim's canonical history") + return refreshed + except (_WorkflowClaimDeferred, LocalActivityExecutionAborted): + raise + except ServerError as error: + if error.reason() == "cancellation_requested": + await self._renew_local_workflow_lease(task) + raise LocalActivityExecutionAborted("prepared operation has an unknown or refused outcome") from error + except Exception as error: + raise LocalActivityExecutionAborted( + "prepared operation could not validate its original authority", + ) from error def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str, Any]: storage = self.external_storage @@ -1987,7 +2187,7 @@ def execute_local(command: RecordLocalActivity) -> Any: cls, task, history, start_input, payload_codec=codec, execute_local=execute_local, ) except _WorkflowClaimDeferred: - log.info("workflow task %s parked until child cleanup finishes", task_id) + log.info("workflow task %s released for durable successor work", task_id) return None except LocalActivityExecutionAborted as e: log.warning("abandoning workflow task %s before local activity commit: %s", task_id, e) @@ -2956,7 +3156,10 @@ def _admit_workflow_work( # The callback owns the reservation so cancellation releases capacity # even when the dispatch coroutine is cancelled before it first runs. - dispatched.add_done_callback(lambda _: self._release_workflow_capacity()) + dispatched.add_done_callback(lambda _: ( + self._release_workflow_capacity() + if task_kind != "workflow" or task.get("task_id") not in self._abandoned_prepared_local_claims else None + )) return dispatched async def _poll_workflow_tasks(self) -> None: @@ -3834,7 +4037,7 @@ async def _shutdown(self) -> None: log.warning("cancelled %d task(s) after shutdown timeout", len(pending)) await asyncio.sleep(0) still_running = {task for task in pending if not task.done()} - if still_running and self._remote_activity_processes: + if still_running and (self._remote_activity_processes or self._prepared_local_activity_processes): # Application drain has ended. Allow bounded process reaping and # stop-report transport without granting more callback work or # extending the run's original cancellation deadline. @@ -3852,9 +4055,10 @@ async def _shutdown(self) -> None: if self._remote_activity_executor is not None: self._remote_activity_executor.shutdown(wait=False, cancel_futures=True) - if self._remote_activity_processes: + if self._remote_activity_processes or self._prepared_local_activity_processes: + kind = "remote" if self._remote_activity_processes else "prepared local" raise RuntimeError( - "worker shutdown has unconfirmed remote callback stop(s); " + f"worker shutdown has unconfirmed {kind} callback stop(s); " "the worker registration remains active" ) diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index a186e57..103a11e 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -631,6 +631,36 @@ def to_server_command( return command +@dataclass(frozen=True) +class PreparedLocalActivityCall: + """An authored local call awaiting admission, with no application execution.""" + + command: RecordLocalActivity + sequence: int + recover: bool = False + cleanup: Mapping[str, str] | None = None + + def descriptor(self, payload_codec: str) -> dict[str, Any]: + command = self.command + descriptor: dict[str, Any] = { + "type": "record_local_activity", "activity_type": command.activity_type, + "execution_mode": "local", "arguments": serializer.envelope(command.arguments, codec=payload_codec), + "payload_codec": payload_codec, + } + if command.retry_policy is not None: + descriptor["retry_policy"] = dict(command.retry_policy) + for name in ("start_to_close_timeout", "schedule_to_close_timeout", "heartbeat_timeout"): + value = getattr(command, name) + if value is not None: + descriptor[name] = value + if self.cleanup is not None: + descriptor["cancellation_cleanup"] = { + "request_id": self.cleanup["request_id"], + "delivery_history_event_id": self.cleanup["delivery_history_event_id"], + } + return descriptor + + class LocalActivityExecutionAborted(Exception): """The workflow task lease could not be trusted after local execution began.""" @@ -2282,6 +2312,7 @@ class ReplayOutcome: message_stream_cursors: list[dict[str, Any]] = field(default_factory=list) message_stream_waits: list[dict[str, Any]] = field(default_factory=list) cancellation_delivery: CancellationDelivery | None = None + prepared_local_activity: PreparedLocalActivityCall | None = None class Replayer: @@ -2655,6 +2686,7 @@ def replay( cancel_requested: bool = False, cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, + prepare_local_activities: bool = False, ) -> ReplayOutcome: return _replay_state( workflow_cls, @@ -2669,6 +2701,7 @@ def replay( cancel_requested=cancel_requested, cancellation_request=cancellation_request, local_activity_executor=local_activity_executor, + prepare_local_activities=prepare_local_activities, ).outcome @@ -3668,6 +3701,7 @@ def _replay_state( cancel_requested: bool = False, cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, + prepare_local_activities: bool = False, stop_at_uncommitted_cancellation: bool = False, ) -> _ReplayState: if payload_codec is not None and payload_codec != serializer.AVRO_CODEC: @@ -3811,11 +3845,14 @@ def _replay_state( ): ctx._accept_message_stream(arguments) - def _state(commands: list[Command]) -> _ReplayState: + def _state( + commands: list[Command], prepared_local_activity: PreparedLocalActivityCall | None = None, + ) -> _ReplayState: return _ReplayState( outcome=ReplayOutcome( commands=commands, cancellation_delivery=cancellation_intent, + prepared_local_activity=prepared_local_activity, message_stream_cursors=[ {"stream_name": name, "through_position": position} for name, position in sorted(ctx._message_stream_cursors.items()) @@ -4018,7 +4055,7 @@ def _assert_next_step_matches(command: Any, offset: int = 0) -> None: if step_index >= len(recorded_steps): return if ( - cancellation.request is not None + (cancellation.request is not None or prepare_local_activities) and recorded_steps[step_index].workflow_sequence != current_call_sequence + offset ): raise NonDeterministicReplayError( @@ -5176,7 +5213,7 @@ def _consume_terminal_condition_reopens() -> None: def _cancellation_boundary(command: Any) -> CancellationDelivery | None: nonlocal authored_sequence, current_call_sequence - if cancellation.request is None: + if cancellation.request is None and not prepare_local_activities: return None current_call_sequence = authored_sequence kind: str | None = None @@ -5223,6 +5260,15 @@ def _cancellation_boundary(command: Any) -> CancellationDelivery | None: operation_sequence, operation_span, ) + def _contains_local_activity(operation: Any) -> bool: + if isinstance(operation, RecordLocalActivity): + return True + if isinstance(operation, list): + return any(_contains_local_activity(member) for member in operation) + if isinstance(operation, SelectGroup): + return any(_contains_local_activity(member) for _, member in operation.operations) + return False + def _assert_cancellation_call_matches(command: Any, boundary: CancellationDelivery) -> None: if isinstance(command, list): leaves, _ = _annotate_parallel_commands(command, boundary.sequence) @@ -5316,6 +5362,10 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl return _terminal_state(stop.value, include_pending=True) first = False _apply_due_receivers() + if prepare_local_activities and isinstance(cmd, list | SelectGroup) and _contains_local_activity(cmd): + raise LocalActivityExecutionAborted( + "prepared local groups require an implemented atomic group consumer", + ) boundary = _cancellation_boundary(cmd) if cancellation.delivery is not None and not cancellation_consumed: cancellation_marker = cancellation.delivery @@ -5768,6 +5818,38 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl continue ctx.logger._set_replaying(False) _assert_pending_step_matches(cmd) + if isinstance(cmd, RecordLocalActivity) and prepare_local_activities: + cleanup: dict[str, str] | None = None + if cancellation.request is not None: + delivery = cancellation.delivery + context = cancellation.request.context + delivery_id = ( + events[cancellation.delivery_index].get("id") + if cancellation.delivery_index is not None else None + ) + if ( + not cancellation_consumed or ctx._cancellation_shield_depth < 1 + or delivery is None or context is None + or current_call_sequence < delivery.sequence + delivery.sequence_span + or not isinstance(delivery_id, str) or not delivery_id.strip() + ): + raise LocalActivityExecutionAborted( + "prepared cleanup requires a shield after canonical cancellation delivery", + ) + cleanup = { + "request_id": context.request_id, "root_request_id": context.root_request_id, + "delivery_history_event_id": delivery_id, + "cleanup_deadline_at": context.to_dict()["cleanup_deadline_at"], + } + started = False + for event in events: + if _workflow_sequence(event.get("payload") or {}) != current_call_sequence: + continue + if _history_event_type(event) == "ActivityStarted": + started = True + elif _history_event_type(event) == "ActivityRetryScheduled": + started = False + return _state(pending, PreparedLocalActivityCall(cmd, current_call_sequence, started, cleanup)) if isinstance(cmd, RecordLocalActivity) and local_activity_executor is not None: local_result = local_activity_executor(cmd) if cmd.outcome is None or cmd.arguments_envelope is None: diff --git a/tests/fixtures/replay_regressions/prepared-local-cold-results.json b/tests/fixtures/replay_regressions/prepared-local-cold-results.json new file mode 100644 index 0000000..3ad3b83 --- /dev/null +++ b/tests/fixtures/replay_regressions/prepared-local-cold-results.json @@ -0,0 +1,23 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "prepared-local-cold-results", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "tests.replay.prepared-local-cold-results", + "input": [], + "payload_codec": "avro" + }, + "history": [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local"}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "activity_execution_id": "first", "activity_attempt_id": "first-attempt"}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}}}, + {"event_type": "ActivityScheduled", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local"}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "activity_execution_id": "second", "activity_attempt_id": "second-attempt"}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}}} + ], + "expected": { + "command_sequence": [{"type": "complete_workflow", "command_type": "CompleteWorkflow", "result": ["fast-value", "fast-value"]}] + } +} diff --git a/tests/integration/prepared_worker.py b/tests/integration/prepared_worker.py new file mode 100644 index 0000000..fb8cbdd --- /dev/null +++ b/tests/integration/prepared_worker.py @@ -0,0 +1,26 @@ +"""A killable owner for prepared callback and canonical cleanup qualification.""" + +from __future__ import annotations + +import asyncio +import os +import sys +from pathlib import Path + +from tests.integration.test_prepared_local_activity import TrackingPreparedClient, prepared_worker + + +async def main() -> None: + queue, owner, mode, trace = sys.argv[1:] + if mode not in {"hold", "finish"}: + raise ValueError("unknown prepared qualification mode") + os.environ["DW_PREPARED_FIXTURE_MODE"] = mode + async with TrackingPreparedClient( + os.environ["DURABLE_WORKFLOW_SERVER_URL"], + token=os.environ.get("DURABLE_WORKFLOW_AUTH_TOKEN", "test-token"), namespace="default", trace_path=Path(trace), + ) as client: + await prepared_worker(client, queue, worker_id=owner).run() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py new file mode 100644 index 0000000..9ac0941 --- /dev/null +++ b/tests/integration/test_prepared_local_activity.py @@ -0,0 +1,215 @@ +"""Opt-in exact-source proof of Python's durable sequential local consumer.""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys +import uuid +from datetime import datetime +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import Client, Worker, activity, workflow +from durable_workflow.errors import WorkflowCancelled +from durable_workflow.workflow import replay +from tests.integration.test_cooperative_cancellation import callback_gone, events, poll_claim, remote_marker + +pytestmark = pytest.mark.usefixtures("prepared_runtime") + + +@pytest.fixture +async def prepared_runtime(server_url: str, server_token: str, monkeypatch: pytest.MonkeyPatch) -> None: + if os.environ.get("DURABLE_WORKFLOW_NATIVE_SOURCE_QUALIFICATION") != "1": + pytest.skip("prepared local consumer requires exact-source Native qualification") + async with Client(server_url, token=server_token, namespace="default") as client: + info = await client.get_cluster_info() + assert info["worker_protocol"]["server_capabilities"]["prepared_local_activities"] is True + assert info["worker_protocol"]["version"] == "1.20" + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + + +@workflow.defn(name="tests.python-prepared-sequential") +class PreparedSequentialWorkflow: + def run(self, ctx: Any, marker: str) -> Any: + yield ctx.upsert_memo({"before": "durable local admission"}) + first = yield ctx.local_activity("tests.python-prepared-local", [marker, "first"], heartbeat_timeout=10) + second = yield ctx.local_activity("tests.python-prepared-local", [marker, "second"], heartbeat_timeout=10) + return [first, second] + + +@activity.defn(name="tests.python-prepared-local") +async def prepared_local(marker: str, phase: str) -> dict[str, Any]: + info = activity.context().info + Path(marker + "." + phase).write_text(json.dumps({ + "callback_pid": os.getpid(), "attempt_id": info.activity_attempt_id, + })) + await activity.context().heartbeat({"phase": phase}) + return {"phase": phase, "bytes": b"\x00\xff", "attempt": info.activity_attempt_id} + + +@workflow.defn(name="tests.python-prepared-cancellation") +class PreparedCancellationWorkflow: + def run(self, ctx: Any, marker: str) -> Any: + try: + yield ctx.local_activity("tests.python-prepared-blocked", [marker, "work", None]) + except WorkflowCancelled as error: + with ctx.cancellation_shield(): + yield ctx.local_activity( + "tests.python-prepared-blocked", [marker, "cleanup", error.request_id], + retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, + ) + return error.request_id + return "not cancelled" + + +@activity.defn(name="tests.python-prepared-blocked") +async def prepared_blocked(marker: str, phase: str, request_id: str | None) -> dict[str, Any]: + info = activity.context().info + path = Path(marker + "." + phase + "." + info.worker_id) + pending = path.with_suffix(path.suffix + ".writing") + pending.write_text(json.dumps({"phase": phase, "callback_pid": os.getpid(), "request_id": request_id, + "activity_attempt_id": info.activity_attempt_id})) + pending.replace(path) + if phase == "work" or os.environ.get("DW_PREPARED_FIXTURE_MODE") == "hold": + # Intentionally no application heartbeat. + await asyncio.Event().wait() + return {"request_id": request_id, "bytes": b"\x00\xff"} + + +def prepared_worker(client: Client, queue: str, **kwargs: Any) -> Worker: + return Worker( + client, task_queue=queue, worker_id=kwargs.pop("worker_id", queue + "-owner"), + workflows=[PreparedSequentialWorkflow, PreparedCancellationWorkflow], + activities=[prepared_local, prepared_blocked], + capabilities=["cooperative_cancellation", "prepared_local_activities"], **kwargs, + ) + + +class TrackingPreparedClient(Client): + def __init__(self, *args: Any, trace_path: Path, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self.trace_path = trace_path + + async def prepared_local_activity_operation(self, **kwargs: Any) -> dict[str, Any]: + receipt = await super().prepared_local_activity_operation(**kwargs) + if kwargs["operation"] != "control" or receipt.get("active") is False: + with self.trace_path.open("a") as stream: + stream.write(json.dumps({"operation": kwargs["operation"], "receipt": receipt}) + "\n") + return receipt + + +async def test_prepared_prefix_two_callbacks_and_cold_replay_use_canonical_history( + server_url: str, server_token: str, tmp_path: Path, +) -> None: + queue = "py-prepared-sequential-" + uuid.uuid4().hex[:8] + marker = str(tmp_path / "callback") + async with Client(server_url, token=server_token, namespace="default") as client: + worker = prepared_worker(client, queue) + await worker._register() + try: + handle = await client.start_workflow(workflow_type="tests.python-prepared-sequential", task_queue=queue, + workflow_id=queue, input=[marker]) + task = await poll_claim(client, worker) + commands = await worker._run_workflow_task(task) + assert commands is not None and [command["type"] for command in commands] == ["complete_workflow"] + result = await handle.result(timeout=10) + assert [item["phase"] for item in result] == ["first", "second"] + assert all(item["bytes"] == b"\x00\xff" for item in result) + assert result[0]["attempt"] != result[1]["attempt"] + history = await events(handle) + kinds = [event["event_type"] for event in history] + assert kinds.count("ActivityStarted") == kinds.count("ActivityCompleted") == 2 + assert kinds.count("ActivityHeartbeatRecorded") == 2 + assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 + outcome = replay(PreparedSequentialWorkflow, history, [marker], run_id=handle.run_id or "", + prepare_local_activities=True) + assert outcome.prepared_local_activity is None and outcome.commands[0].result == result # type: ignore[union-attr] + for phase in ("first", "second"): + proof = json.loads(Path(marker + "." + phase).read_text()) + await callback_gone(proof["callback_pid"]) + print("prepared sequential history: " + json.dumps(history)) + finally: + await worker.stop() + + +async def test_prepared_callback_stop_cleanup_sigkill_and_cold_recovery_keep_original_30_second_deadline( + server_url: str, server_token: str, tmp_path: Path, +) -> None: + queue = "py-prepared-cancel-" + uuid.uuid4().hex[:8] + marker = str(tmp_path / "callback") + trace_path = tmp_path / "receipts.jsonl" + processes: list[asyncio.subprocess.Process] = [] + logs: list[Any] = [] + + async def owner(name: str, mode: str) -> asyncio.subprocess.Process: + log = (tmp_path / (name + ".log")).open("wb") + logs.append(log) + process = await asyncio.create_subprocess_exec( + sys.executable, "-m", "tests.integration.prepared_worker", queue, name, mode, str(trace_path), + stdout=log, stderr=log, + ) + processes.append(process) + return process + + async with Client(server_url, token=server_token, namespace="default") as client: + try: + handle = await client.start_workflow(workflow_type="tests.python-prepared-cancellation", task_queue=queue, + workflow_id=queue, input=[marker]) + first = await owner(queue + "-first", "hold") + work = await remote_marker(Path(marker + ".work." + queue + "-first")) + accepted = await handle.request_cancellation(cleanup_timeout_seconds=30) + original = accepted["cancellation_request"] + requested_history = await events(handle) + original_context = next(event["payload"]["cancellation"] for event in requested_history + if event["event_type"] == "CooperativeCancellationRequested") + cleanup = await remote_marker(Path(marker + ".cleanup." + queue + "-first")) + assert cleanup["request_id"] == original["request_id"] + await callback_gone(work["callback_pid"]) + first.kill() + await asyncio.wait_for(first.wait(), timeout=10) + await callback_gone(cleanup["callback_pid"]) + duplicate = await handle.request_cancellation(cleanup_timeout_seconds=90) + for field in ("request_id", "requested_at", "cleanup_deadline_at"): + assert duplicate["cancellation_request"][field] == original[field] + await owner(queue + "-replacement", "finish") + deadline = datetime.fromisoformat(original["cleanup_deadline_at"].replace("Z", "+00:00")) + timeout = (deadline - datetime.now(deadline.tzinfo)).total_seconds() + assert timeout > 0 + with pytest.raises(WorkflowCancelled): + await handle.result(timeout=timeout) + history = await events(handle) + kinds = [event["event_type"] for event in history] + assert kinds.count("CooperativeCancellationDelivered") == kinds.count("WorkflowCancelled") == 1 + assert kinds.count("ActivityCancellationAcknowledged") == 1 + assert "ActivityHeartbeatRecorded" not in kinds + terminal = next(event for event in history if event["event_type"] == "WorkflowCancelled") + assert datetime.fromisoformat(terminal["timestamp"].replace("Z", "+00:00")) < deadline + receipts = [json.loads(line) for line in trace_path.read_text().splitlines()] + admissions = [entry["receipt"] for entry in receipts + if entry["operation"] == "prepare" and entry["receipt"].get("cancellation_cleanup")] + assert len(admissions) == 2 + assert admissions[0]["activity_attempt_id"] != admissions[1]["activity_attempt_id"] + assert admissions[0]["cancellation_cleanup"] == admissions[1]["cancellation_cleanup"] + authority = admissions[0]["cancellation_cleanup"] + assert authority["request_id"] == original["request_id"] + assert authority["root_request_id"] == original_context["root_request_id"] + assert datetime.fromisoformat(authority["cleanup_deadline_at"].replace("Z", "+00:00")) == deadline + recoveries = [entry["receipt"] for entry in receipts if entry["operation"] == "recover"] + assert len(recoveries) == 1 and recoveries[0]["callback_stop_state"] == "unknown" + replacement = await remote_marker(Path(marker + ".cleanup." + queue + "-replacement")) + await callback_gone(replacement["callback_pid"]) + print("prepared cleanup recovery receipts: " + json.dumps(receipts)) + print("prepared cleanup recovery history: " + json.dumps(history)) + finally: + for process in processes: + if process.returncode is None: + process.kill() + await asyncio.wait_for(process.wait(), timeout=10) + for log in logs: + log.close() + for path in tmp_path.glob("*.log"): + print(path.name + ": " + path.read_text()) diff --git a/tests/test_prepared_local_activity_worker.py b/tests/test_prepared_local_activity_worker.py new file mode 100644 index 0000000..1e49e0b --- /dev/null +++ b/tests/test_prepared_local_activity_worker.py @@ -0,0 +1,483 @@ +from __future__ import annotations + +import asyncio +import json +import os +import shutil +import time +from copy import deepcopy +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import activity, serializer, workflow +from durable_workflow._prepared_local_activity import PreparedAttempt, PreparedCancellationObserved +from durable_workflow.client import Client +from durable_workflow.errors import WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted, replay +from tests.test_activity_process import blocked_without_python_progress, wait_for_exit, wait_for_file, wait_for_release +from tests.test_cancellation_context import delivery, request, snapshot +from tests.test_replay_regression_corpus import PreparedLocalColdResultsWorkflow +from tests.test_worker import compatible_cluster_info + + +@workflow.defn(name="prepared.sequential") +class SequentialWorkflow: + def run(self, ctx: workflow.WorkflowContext, marker: str, prefix: bool = False): # type: ignore[no-untyped-def] + if prefix: + yield ctx.upsert_memo({"before": "local"}) + return (yield ctx.local_activity("prepared.callback", [marker])) + + +@workflow.defn(name="prepared.cleanup") +class CleanupWorkflow: + def run(self, ctx: workflow.WorkflowContext, shield: bool = True): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled: + if shield: + with ctx.cancellation_shield(): + yield ctx.local_activity("prepared.callback", []) + else: + yield ctx.local_activity("prepared.callback", []) + return "cleaned" + + +@workflow.defn(name="prepared.group") +class GroupWorkflow: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + return (yield [ctx.local_activity("prepared.callback", []), ctx.start_timer(1)]) + + +@activity.defn(name="prepared.callback") +def typed_callback(marker: str) -> dict[str, Any]: + Path(marker).write_text(str(os.getpid())) + return {"value": b"\x00\xff", "attempt": activity.context().info.activity_attempt_id} + + +def timestamp(seconds: float = 0) -> str: + return (datetime.now(timezone.utc) + timedelta(seconds=seconds)).isoformat(timespec="microseconds").replace( + "+00:00", "Z", + ) + + +def admission(**changes: Any) -> dict[str, Any]: + return { + "prepared": True, "duplicate": False, "reason": None, + "workflow_task_id": "task", "workflow_task_attempt": 3, "lease_owner": "owner", + "worker_attempt_id": "nonce", "activity_execution_id": "execution", "activity_attempt_id": "attempt", + "attempt_number": 1, "server_time": timestamp(), "lease_expires_at": timestamp(30), + "start_to_close_deadline_at": None, "schedule_to_close_deadline_at": None, "heartbeat_deadline_at": None, + "cancellation_cleanup": None, **changes, + } + + +def admitted(receipt: dict[str, Any], **changes: Any) -> PreparedAttempt: + return PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=None, cleanup=None, request_started=time.monotonic(), **changes, + ) + + +def control(receipt: dict[str, Any], **changes: Any) -> dict[str, Any]: + return { + **receipt, "active": True, "renewed": True, "stop_required": False, + "workflow_lease_expires_at": receipt["lease_expires_at"], + "heartbeat_recorded": False, "heartbeat_history_event_id": None, **changes, + } + + +@pytest.mark.parametrize("changes", [ + {"prepared": False}, {"duplicate": 1}, {"reason": "refused"}, {"workflow_task_id": "other"}, + {"workflow_task_attempt": True}, {"workflow_task_attempt": 4}, {"lease_owner": "other"}, + {"worker_attempt_id": "other"}, {"attempt_number": True}, {"attempt_number": 0}, + {"activity_attempt_id": ""}, {"activity_execution_id": " "}, {"server_time": "2026-02-30T00:00:00Z"}, + {"lease_expires_at": "2026-01-01T00:00:00Z"}, {"heartbeat_deadline_at": "2027-01-01T00:00:00Z"}, + {"cancellation_cleanup": {"request_id": "invented"}}, +]) +def test_admission_refuses_changed_identity_and_invented_authority(changes: dict[str, Any]) -> None: + with pytest.raises(LocalActivityExecutionAborted): + admitted(admission(**changes)) + + +@pytest.mark.parametrize("field", ["reason", "start_to_close_deadline_at", "schedule_to_close_deadline_at", + "heartbeat_deadline_at"]) +def test_admission_requires_explicit_receipt_fields(field: str) -> None: + receipt = admission() + del receipt[field] + with pytest.raises(LocalActivityExecutionAborted): + admitted(receipt) + + +@pytest.mark.parametrize("changes", [ + {"active": 1}, {"renewed": False}, {"stop_required": True}, {"workflow_task_attempt": True}, + {"activity_attempt_id": "other"}, {"lease_owner": "other"}, {"reason": "refused"}, + {"heartbeat_recorded": True, "heartbeat_history_event_id": "invented-progress"}, + {"heartbeat_deadline_at": "2027-01-01T00:00:00Z"}, + {"start_to_close_deadline_at": "2027-01-01T00:00:00Z"}, + {"workflow_lease_expires_at": "2026-01-01T00:00:00Z"}, + {"cancellation_cleanup": {"request_id": "invented"}}, +]) +def test_supervisor_control_cannot_change_authority_or_report_application_progress(changes: dict[str, Any]) -> None: + receipt = admission() + attempt = admitted(receipt) + with pytest.raises(LocalActivityExecutionAborted): + attempt.validate_control(control(receipt, **changes)) + + +def test_application_heartbeat_alone_advances_its_deadline() -> None: + receipt = admission(heartbeat_deadline_at=timestamp(9)) + attempt = PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=10, cleanup=None, request_started=time.monotonic(), + ) + refreshed = control(receipt, renewed=False, heartbeat_recorded=True, heartbeat_history_event_id="heartbeat", + server_time=timestamp(1), heartbeat_deadline_at=timestamp(10)) + attempt.validate_control(refreshed, heartbeat=True) + assert attempt.deadlines["heartbeat_deadline_at"] == refreshed["heartbeat_deadline_at"] + attempt.validate_control(control(refreshed, renewed=True, heartbeat_recorded=False, + heartbeat_history_event_id=None)) + + +def test_transport_elapsed_time_cannot_restart_an_authority_budget() -> None: + receipt = admission(lease_expires_at=timestamp(2)) + with pytest.raises(LocalActivityExecutionAborted, match="budget expired"): + PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=None, cleanup=None, request_started=time.monotonic() - 3, + ) + + +@pytest.mark.parametrize("field", ["request_id", "root_request_id", "delivery_history_event_id", "cleanup_deadline_at"]) +def test_cleanup_admission_and_control_cannot_replace_original_cascade_authority(field: str) -> None: + cleanup = {"request_id": "request-1", "root_request_id": "root-1", "delivery_history_event_id": "delivery-1", + "cleanup_deadline_at": timestamp(20)} + receipt = admission(cancellation_cleanup=cleanup, heartbeat_deadline_at=cleanup["cleanup_deadline_at"], + start_to_close_deadline_at=cleanup["cleanup_deadline_at"], + schedule_to_close_deadline_at=cleanup["cleanup_deadline_at"]) + attempt = PreparedAttempt.admitted( + receipt, task_id="task", run_id="run-1", owner="owner", epoch=3, nonce="nonce", + heartbeat_timeout=None, cleanup=cleanup, request_started=time.monotonic(), + ) + changed = {**cleanup, field: timestamp(30) if field == "cleanup_deadline_at" else "replacement"} + with pytest.raises(LocalActivityExecutionAborted, match="cleanup authority"): + attempt.validate_control(control(receipt, cancellation_cleanup=changed)) + with pytest.raises(LocalActivityExecutionAborted, match="cleanup authority"): + PreparedAttempt.admitted( + {**receipt, "cancellation_cleanup": changed}, task_id="task", run_id="run-1", owner="owner", + epoch=3, nonce="nonce", heartbeat_timeout=None, cleanup=cleanup, request_started=time.monotonic(), + ) + + +def test_replay_captures_a_fresh_local_call_before_invoking_its_executor() -> None: + outcome = replay(SequentialWorkflow, [], ["unused", True], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("callback ran before admission")) + assert [type(command).__name__ for command in outcome.commands] == ["UpsertMemo"] + call = outcome.prepared_local_activity + assert call is not None and call.sequence == 2 and call.recover is False and call.cleanup is None + descriptor = call.descriptor("avro") + assert "outcome" not in descriptor and "result" not in descriptor + assert serializer.decode_envelope(descriptor["arguments"]) == ["unused"] + + +@pytest.mark.parametrize("last_event,recover", [("ActivityStarted", True), ("ActivityRetryScheduled", False)]) +def test_cold_replay_distinguishes_unknown_started_callback_from_durable_retry(last_event: str, recover: bool) -> None: + history = [{"event_type": "ActivityScheduled", "payload": { + "sequence": 1, "activity_type": "prepared.callback", "execution_mode": "local", + }}, {"event_type": last_event, "payload": { + "sequence": 1, "activity_type": "prepared.callback", "execution_mode": "local", + }}] + outcome = replay(SequentialWorkflow, history, ["unused"], prepare_local_activities=True) + assert outcome.commands == [] + assert outcome.prepared_local_activity is not None + assert outcome.prepared_local_activity.sequence == 1 and outcome.prepared_local_activity.recover is recover + + +def test_immutable_cold_results_skip_completed_callback_and_recover_unfinished_original_attempt() -> None: + fixture = json.loads(( + Path(__file__).parent / "fixtures/replay_regressions/prepared-local-cold-results.json" + ).read_text()) + history = fixture["history"] + outcome = replay(PreparedLocalColdResultsWorkflow, history[:-1], [], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("unfinished Started callback was reexecuted")) + call = outcome.prepared_local_activity + assert call is not None and call.sequence == 2 and call.recover is True + assert call.command.activity_type == "prepared.second" and outcome.commands == [] + completed = replay(PreparedLocalColdResultsWorkflow, history, [], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("completed callback was reexecuted")) + assert completed.prepared_local_activity is None + assert completed.commands[0].result == ["fast-value", "fast-value"] # type: ignore[union-attr] + + +def test_cleanup_capture_preserves_the_original_delivery_root_and_deadline() -> None: + marker = {**delivery(), "id": "canonical-delivery"} + original = [request(), marker] + outcome = replay(CleanupWorkflow, original, [], run_id="run-1", prepare_local_activities=True) + call = outcome.prepared_local_activity + assert call is not None and call.sequence == 2 + assert call.cleanup == { + "request_id": "request-1", "root_request_id": "root-1", "delivery_history_event_id": "canonical-delivery", + "cleanup_deadline_at": snapshot()["cleanup_deadline_at"], + } + assert call.descriptor("avro")["cancellation_cleanup"] == { + "request_id": "request-1", "delivery_history_event_id": "canonical-delivery", + } + original[1]["id"] = "" + with pytest.raises(LocalActivityExecutionAborted, match="canonical cancellation delivery"): + replay(CleanupWorkflow, original, [], run_id="run-1", prepare_local_activities=True) + with pytest.raises(LocalActivityExecutionAborted, match="shield"): + replay(CleanupWorkflow, [request(), marker], [False], run_id="run-1", prepare_local_activities=True) + + +def test_sequential_consumer_refuses_a_local_parallel_group_before_any_callback() -> None: + with pytest.raises(LocalActivityExecutionAborted, match="atomic group consumer"): + replay(GroupWorkflow, [], [], prepare_local_activities=True, + local_activity_executor=lambda _: pytest.fail("group callback ran before complete admission")) + + +class PreparedServer: + def __init__(self) -> None: + self.client = AsyncMock(spec=Client) + self.history: list[dict[str, Any]] = [] + self.trace: list[str] = [] + self.receipt: dict[str, Any] = {} + self.stop = False + self.bad_admission = False + self.omit_outcome_history = False + self.fail_outcome = False + self.marker: Path | None = None + self.client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": { + "prepared_local_activities": True, "cooperative_cancellation": True, + "workflow_memo_updates": {"supported": True}, + "supported_workflow_task_commands": ["upsert_memo"], + }, + }) + self.client.register_worker.return_value = {"registered": True} + self.client.prepared_local_activity_operation.side_effect = self.operation + self.client.workflow_task_history.side_effect = self.page + self.client.complete_workflow_task.return_value = {"outcome": "completed"} + + async def page(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs == {"task_id": "task", "lease_owner": "owner", "workflow_task_attempt": 3, + "next_history_page_token": "canonical-page"} + self.trace.append("history") + return {"history_events": deepcopy(self.history), "next_history_page_token": None} + + async def operation(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs["task_id"] == "task" and kwargs["lease_owner"] == "owner" + assert kwargs["workflow_task_attempt"] == 3 + name = kwargs["operation"] + self.trace.append(name) + body = kwargs["body"] + if name == "prepare": + assert "outcome" not in body["descriptor"] + self.receipt = admission(worker_attempt_id=body["worker_attempt_id"]) + if self.bad_admission: + self.receipt["activity_attempt_id"] = "" + payload = {"sequence": body["sequence"], "activity_type": "prepared.callback", "execution_mode": "local", + "activity_execution_id": "execution", "activity_attempt_id": "attempt"} + self.history.extend([{"id": "scheduled", "event_type": "ActivityScheduled", "payload": payload}, + {"id": "started", "event_type": "ActivityStarted", "payload": payload}]) + return self.receipt + assert kwargs["activity_attempt_id"] == "attempt" + if name == "control": + assert body == {"renew_lease": True} + if not self.stop: + return control(self.receipt) + context = snapshot() + context.update(requested_at=timestamp(), cleanup_deadline_at=timestamp(30)) + return control(self.receipt, active=False, stop_required=True, renewed=False, + reason="cancellation_requested", fenced=True, cancellation_request=context, + cancellation_history_event_id="cancelled", history_refresh_page_token="canonical-page") + if name == "acknowledge-cancellation": + assert body == {"request_id": "request-1"} + assert self.marker is not None + pid = int(self.marker.read_text()) + with pytest.raises(ProcessLookupError): + os.kill(pid, 0) + return {"acknowledged": True, "duplicate": False, "reason": None, "history_event_id": "joined"} + assert name == "outcome" + if self.fail_outcome: + raise TimeoutError("outcome acknowledgment lost") + if not self.omit_outcome_history: + payload = {**self.history[-1]["payload"], "result": body["report"]["result"]} + self.history.append({"id": "completed", "event_type": "ActivityCompleted", "payload": payload}) + return {**self.receipt, "recorded": True, "event_id": "completed", "event_type": "ActivityCompleted", + "workflow_run_id": "run-1", "recorded_at": timestamp(), "claim_released": False, + "created_task_ids": [], "history_refresh_page_token": "canonical-page"} + + async def worker(self, monkeypatch: pytest.MonkeyPatch, *, handler: Any = typed_callback) -> Worker: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + worker = Worker(self.client, task_queue="queue", worker_id="owner", workflows=[SequentialWorkflow], + capabilities=["cooperative_cancellation", "prepared_local_activities"]) + await worker._register() + worker.activities["prepared.callback"] = handler + return worker + + def task(self, marker: Path) -> dict[str, Any]: + self.marker = marker + return {"task_id": "task", "workflow_type": "prepared.sequential", "workflow_id": "workflow-1", + "run_id": "run-1", "workflow_task_attempt": 3, "payload_codec": "avro", + "arguments": serializer.envelope([str(marker)]), "history_events": deepcopy(self.history)} + + +async def test_worker_runs_only_an_admitted_process_then_replays_canonical_result( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch) + commands = await worker._run_workflow_task(server.task(tmp_path / "callback")) + assert commands is not None and [command["type"] for command in commands] == ["complete_workflow"] + assert serializer.decode_envelope(commands[0]["result"]) == {"value": b"\x00\xff", "attempt": "attempt"} + assert server.trace == ["prepare", "control", "control", "outcome", "history"] + assert not worker._prepared_local_activity_processes + server.client.fail_workflow_task.assert_not_awaited() + await wait_for_exit(int((tmp_path / "callback").read_text())) + + +async def test_bad_admission_never_spawns_and_never_completes_or_fails_a_claim( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + server.bad_admission = True + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(server.task(tmp_path / "callback")) is None + assert not (tmp_path / "callback").exists() + assert server.trace == ["prepare"] + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("failure", ["fail_outcome", "omit_outcome_history"]) +async def test_lost_or_noncanonical_outcome_abandons_without_reexecuting_callback( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, failure: str, +) -> None: + server = PreparedServer() + setattr(server, failure, True) + worker = await server.worker(monkeypatch) + assert await worker._run_workflow_task(server.task(tmp_path / "callback")) is None + assert server.trace.count("prepare") == 1 and server.trace.count("outcome") == 1 + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + assert not worker._prepared_local_activity_processes + + +async def test_control_stops_and_joins_a_gil_blocked_local_callback_without_application_heartbeats( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch, handler=blocked_without_python_progress) + marker = tmp_path / "callback" + task = server.task(marker) + outcome = replay(SequentialWorkflow, [], [str(marker)], prepare_local_activities=True) + pending = asyncio.create_task(worker._execute_prepared_local_activity(task, [], outcome)) + try: + await wait_for_file(marker) + pid = int(marker.read_text()) + started = time.monotonic() + server.stop = True + with pytest.raises(PreparedCancellationObserved): + await asyncio.wait_for(pending, timeout=7.0) + assert time.monotonic() - started < 4.0 + await wait_for_exit(pid) + assert server.trace.count("acknowledge-cancellation") == 1 + assert "heartbeat" not in server.trace and "outcome" not in server.trace + assert task["_prepared_cancellation_context"]["request_id"] == "request-1" + assert not worker._prepared_local_activity_processes + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + finally: + if not pending.done(): + pending.cancel() + await asyncio.gather(pending, return_exceptions=True) + + +async def test_prepared_capability_requires_actual_bridge_and_never_advertises_group_support( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch) + manifest = server.client.register_worker.await_args.kwargs["capability_manifest"] + assert manifest["prepared_local_activities"]["implementation"] == "durable_sequential_admission" + assert "prepared_local_activity_groups" not in manifest + info = server.client.get_cluster_info.return_value + del info["worker_protocol"]["server_capabilities"]["prepared_local_activities"] + with pytest.raises(RuntimeError, match="installed admission bridge"): + await worker._register() + with pytest.raises(ValueError, match="no prepared local group consumer"): + Worker(server.client, task_queue="queue", capabilities=["prepared_local_activity_groups"]) + + +async def test_descendant_stop_keeps_local_observation_time_and_canonical_root_time_distinct( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch) + task = server.task(tmp_path / "unused") + root = snapshot() + observed = {"request_id": root["request_id"], "requested_at": "2026-10-01T00:00:05Z", + "cleanup_deadline_at": root["cleanup_deadline_at"], "history_refresh_page_token": "canonical-page"} + server.client.heartbeat_workflow_task.return_value = { + "task_id": "task", "lease_owner": "owner", "workflow_task_attempt": 3, "renewed": True, + "cancellation_request": observed, + } + + async def stopped(*_: Any) -> list[dict[str, Any]]: + task["_prepared_cancellation_context"] = root + server.history = [request()] + raise PreparedCancellationObserved("original callback joined") + + async def delivered(**kwargs: Any) -> dict[str, Any]: + assert kwargs["call_kind"] == "local_activity" + server.history.append({**delivery(), "payload": {**delivery()["payload"], "call_kind": "local_activity"}}) + return {"delivered": True} + + worker._execute_prepared_local_activity = stopped # type: ignore[method-assign] + server.client.deliver_workflow_cancellation.side_effect = delivered + outcome, _ = await worker._replay_workflow_claim( + SequentialWorkflow, task, [], ["unused"], payload_codec="avro", + execute_local=lambda _: pytest.fail("callback reexecuted after joined stop"), + ) + assert outcome.prepared_local_activity is None and len(outcome.commands) == 1 + assert task["cancellation_request"] == observed + assert task["cancellation_request"]["requested_at"] != root["requested_at"] + assert server.client.deliver_workflow_cancellation.await_count == 1 + + +async def test_lost_prepared_supervisor_retains_workflow_capacity_and_never_reports_stop( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch, handler=wait_for_release) + worker.max_concurrent_workflow_tasks = 1 + worker._wf_semaphore = asyncio.Semaphore(1) + marker = tmp_path / "callback" + await worker._reserve_workflow_capacity() + pending = worker._admit_workflow_work(server.task(marker), "workflow") + await wait_for_file(marker) + callback = next(iter(worker._prepared_local_activity_processes)) + pid = int(marker.read_text()) + try: + callback._supervisor.kill() + assert await asyncio.wait_for(pending, timeout=4.0) is None + os.kill(pid, 0) + assert callback.stopped is False + assert worker._current_task_slots()["workflow_available"] == 0 and worker._workflow_reserved == 1 + assert "acknowledge-cancellation" not in server.trace and "outcome" not in server.trace + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + with pytest.raises(RuntimeError, match="unconfirmed prepared local callback stop.*registration remains active"): + await worker._shutdown() + server.client.deregister_worker_registration.assert_not_awaited() + finally: + Path(str(marker) + ".release").touch() + await wait_for_exit(pid) + shutil.rmtree(callback.directory, ignore_errors=True) + worker._prepared_local_activity_processes.discard(callback) + worker._abandoned_prepared_local_claims.discard("task") + worker._release_workflow_capacity() diff --git a/tests/test_replay_regression_corpus.py b/tests/test_replay_regression_corpus.py index 972cebd..64e1f22 100644 --- a/tests/test_replay_regression_corpus.py +++ b/tests/test_replay_regression_corpus.py @@ -212,6 +212,14 @@ def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] return (yield ctx.local_activity("golden.local", [])) +@workflow.defn(name="tests.replay.prepared-local-cold-results") +class PreparedLocalColdResultsWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + first = yield ctx.local_activity("prepared.first", []) + second = yield ctx.local_activity("prepared.second", []) + return [first, second] + + @workflow.defn(name="tests.replay.cooperative-reopened-condition-cleanup") class CooperativeReopenedConditionCleanupWorkflow: def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] @@ -240,6 +248,7 @@ def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] GoldenTimeoutWaitWorkflow, GoldenVersionMarkerWorkflow, LocalActivityColdResultWorkflow, + PreparedLocalColdResultsWorkflow, MessageStreamConsumerWorkflow, NestedParallelPathWorkflow, ParallelMetadataProducerWorkflow, From 0a7c68b140d4b74b73120716b2dff97773ae58ae Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 12:09:25 +0000 Subject: [PATCH 25/41] Consume atomic prepared local activity groups with joined cancellation --- docs/cooperative-cancellation-design.md | 29 +- src/durable_workflow/_activity_process.py | 14 + .../_prepared_local_activity.py | 31 +- src/durable_workflow/worker.py | 296 ++++++++++--- src/durable_workflow/workflow.py | 171 ++++++-- .../prepared-local-group-cold-results.json | 23 ++ .../test_prepared_local_activity.py | 128 ++++-- tests/test_activity_process.py | 34 ++ tests/test_prepared_local_activity_groups.py | 390 ++++++++++++++++++ tests/test_prepared_local_activity_worker.py | 2 +- tests/test_replay_regression_corpus.py | 7 + 11 files changed, 998 insertions(+), 127 deletions(-) create mode 100644 tests/fixtures/replay_regressions/prepared-local-group-cold-results.json create mode 100644 tests/test_prepared_local_activity_groups.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index 832a10e..a529ecb 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -118,8 +118,10 @@ effects already performed. The source candidate can explicitly request both `cooperative_cancellation` and `prepared_local_activities` in Worker capabilities. Registration requires source protocol 1.20 and actual Server discovery of its installed admission bridge. -The manifest advertises `durable_sequential_admission`. Python refuses prepared -local parallel and selection groups until its atomic group consumer exists. +The manifest advertises `durable_sequential_admission`. Ordinary parallel groups +add explicit `prepared_local_activity_groups` capability and require the Server's +installed atomic admission bridge. Their manifest advertises +`durable_atomic_all_admission`. Selection and turn-closing waits remain refused. Replay captures the authored local call and sequence before application code runs. Earlier side effects, version markers and metadata commands obtain a @@ -146,6 +148,29 @@ control preserve its original local request ID, root ID, delivery history event ID and deadline. No replacement or duplicate request grants a fresh cleanup budget. Connected exact-source qualification remains separate from publication. +## Prepared local parallel groups + +An ordinary list may contain local activities, remote activities, children, +timers and nested lists, with at most 100 total leaves. The worker checkpoints +the complete authored batch atomically, including every nested position. It +validates all opening history and each canonical local execution identity before +preparing callbacks. Every local admission must validate before any callback +spawns. Earlier metadata and side effects use a separate retained checkpoint. + +Callbacks execute concurrently with independent authority observation. Native +owns outcome history, deadlines, retries and unknown-stop recovery. A retry, +receipt loss or authority loss stops and joins siblings before returning the +claim. Cancellation joins the entire group before workflow delivery or cleanup +replay. A stop acknowledgment proves its own callback physically joined, or +that no callback spawned. An unconfirmed join retains workflow capacity and +worker registration. + +Cold replay preserves completed siblings and recovers only unfinished Started +attempts. A durable retry must release the original claim before new admission. +Results preserve authored nested positions despite settlement order. A cleanup +group requires a shield after canonical delivery and every local member retains +the same original root, delivery event and immutable deadline. + ## Remaining qualification Rust policy parity, portable activity policies, nested scopes and deterministic diff --git a/src/durable_workflow/_activity_process.py b/src/durable_workflow/_activity_process.py index 39f4289..63e8d70 100644 --- a/src/durable_workflow/_activity_process.py +++ b/src/durable_workflow/_activity_process.py @@ -284,6 +284,20 @@ async def wait_for_stop() -> None: raise async def _finish(self) -> None: + # Once stop evidence arrives, cancellation cannot interrupt the OS + # join and leave a physically exited callback falsely unconfirmed. + join = asyncio.create_task(self._confirm_join()) + cancelled = False + while not join.done(): + try: + await asyncio.shield(join) + except asyncio.CancelledError: + cancelled = True + join.result() + if cancelled: + raise asyncio.CancelledError + + async def _confirm_join(self) -> None: await self._send(b"F") await self.close() if self._exitcode != 0: diff --git a/src/durable_workflow/_prepared_local_activity.py b/src/durable_workflow/_prepared_local_activity.py index 1e106f0..856fa65 100644 --- a/src/durable_workflow/_prepared_local_activity.py +++ b/src/durable_workflow/_prepared_local_activity.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import contextlib import re import time from collections.abc import Callable, Mapping @@ -230,6 +231,7 @@ def __init__( self.observe = observe self.shutdown = shutdown self.stop_request: CancellationContext | None = None + self.stop_acknowledged = False self.lock = asyncio.Lock() self.callback: SupervisedCallback | None = None @@ -314,18 +316,26 @@ async def heartbeat(details: dict[str, Any] | None) -> None: self.attempt.validate_outcome(receipt) return receipt finally: - for pending in background: - if not pending.done(): - pending.cancel() - await asyncio.gather(*background, return_exceptions=True) - if not callback.stopped: - await callback.stop() - _require(callback.stopped, "prepared callback stop remains unconfirmed") - processes.discard(callback) - await self.acknowledge_stop() + async def finish() -> None: + for pending in background: + if not pending.done(): + pending.cancel() + await asyncio.gather(*background, return_exceptions=True) + if not callback.stopped: + await callback.stop() + _require(callback.stopped, "prepared callback stop remains unconfirmed") + processes.discard(callback) + await self.acknowledge_stop() + + joined = asyncio.create_task(finish()) + while not joined.done(): + # The original cancellation still propagates after finally. + with contextlib.suppress(asyncio.CancelledError): + await asyncio.shield(joined) + joined.result() async def acknowledge_stop(self) -> None: - if self.stop_request is None: + if self.stop_request is None or self.stop_acknowledged: return # A physical join, or the absence of any spawned callback, is proved # before this diagnostic receipt. It grants no execution authority. @@ -338,3 +348,4 @@ async def acknowledge_stop(self) -> None: and "reason" in receipt and receipt["reason"] is None, "joined prepared callback stop was not acknowledged") _text(receipt.get("history_event_id")) + self.stop_acknowledged = True diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 329e193..b3d871c 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -94,6 +94,7 @@ FailWorkflow, LocalActivityExecutionAborted, NexusServiceCall, + PreparedLocalActivityCall, RecordLocalActivity, RecordSideEffect, ReplayOutcome, @@ -1052,8 +1053,11 @@ def __init__( self.capabilities = tuple(dict.fromkeys(capability.strip() for capability in capabilities)) self._cooperative_cancellation_supported = False self._prepared_local_activities_supported = False - if "prepared_local_activity_groups" in self.capabilities: - raise ValueError("Python has no prepared local group consumer yet") + self._prepared_local_activity_groups_supported = False + if "prepared_local_activity_groups" in self.capabilities and ( + "prepared_local_activities" not in self.capabilities + ): + raise ValueError("prepared local groups require prepared_local_activities capability") if "prepared_local_activities" in self.capabilities and "cooperative_cancellation" not in self.capabilities: raise ValueError("prepared local activities require cooperative_cancellation capability") if any(not capability for capability in self.capabilities): @@ -1243,6 +1247,13 @@ async def _register(self) -> None: raise RuntimeError( "prepared_local_activity_not_supported: Server must advertise its installed admission bridge", ) + self._prepared_local_activity_groups_supported = ( + "prepared_local_activity_groups" in self.capabilities and self._prepared_local_activities_supported + and isinstance(server_capabilities, Mapping) + and server_capabilities.get("prepared_local_activity_groups") is True + ) + if "prepared_local_activity_groups" in self.capabilities and not self._prepared_local_activity_groups_supported: + raise RuntimeError("prepared_local_group_not_supported: Server must advertise its installed atomic bridge") self._validate_cooperative_activity_handlers() self._query_tasks_supported = _server_supports_query_tasks(info) self._workflow_memo_updates_supported = _server_supports_workflow_memo_updates(info) @@ -1298,6 +1309,10 @@ async def _register(self) -> None: "supported": True, "minimum_protocol_version": "1.20", "implementation": "durable_sequential_admission", }} if self._prepared_local_activities_supported else {}), + **({"prepared_local_activity_groups": { + "supported": True, "minimum_protocol_version": "1.20", + "implementation": "durable_atomic_all_admission", + }} if self._prepared_local_activity_groups_supported else {}), }, task_slots=self._current_task_slots(), process_metrics=self._current_process_metrics(), @@ -1551,7 +1566,11 @@ async def _replay_workflow_claim( cancel_requested=bool(task.get("cancel_requested", False)) and state.request is None, cancellation_request=task.get("cancellation_request"), local_activity_executor=execute_local, prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, ) + if outcome.prepared_local_activity_group is not None: + history = await self._execute_prepared_local_activity_group(task, history, outcome) + continue if outcome.prepared_local_activity is not None: history = await self._execute_prepared_local_activity(task, history, outcome) continue @@ -1635,35 +1654,7 @@ async def _execute_prepared_local_activity( codec = _validate_payload_codec(task.get("payload_codec")) or serializer.AVRO_CODEC try: if outcome.commands: - commands = commands_to_server_commands(outcome.commands, self.task_queue, payload_codec=codec) - start = call.sequence - len(commands) - if start < 1 or any(command["type"] not in { - "record_side_effect", "record_version_marker", "upsert_memo", "upsert_search_attributes", - } for command in commands): - raise LocalActivityExecutionAborted("prepared prefix has no supported authored sequence range") - if any(command["type"] == "upsert_memo" for command in commands) and ( - not self._workflow_memo_updates_supported - ): - raise LocalActivityExecutionAborted( - "prepared memo prefix requires negotiated workflow memo updates", - ) - checkpoint = hashlib.sha256(json.dumps( - [task_id, self.worker_id, epoch, start, commands], sort_keys=True, allow_nan=False, - ).encode()).hexdigest() - receipt = await self.client.prepared_local_activity_operation( - task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, operation="checkpoint", - body={"checkpoint_id": checkpoint, "start_sequence": start, "commands": commands}, - ) - if (receipt.get("checkpointed") is not True or type(receipt.get("duplicate")) is not bool - or receipt.get("checkpoint_id") != checkpoint - or receipt.get("task_id") != task_id or receipt.get("workflow_run_id") != task["run_id"] - or type(receipt.get("workflow_task_attempt")) is not int - or receipt["workflow_task_attempt"] != epoch or receipt.get("lease_owner") != self.worker_id - or receipt.get("start_sequence") != start or receipt.get("next_sequence") != call.sequence - or "reason" not in receipt - or receipt["reason"] is not None): - raise LocalActivityExecutionAborted("prepared prefix lacks an original-claim checkpoint receipt") - return await self._refresh_prepared_local_history(task, receipt) + return await self._checkpoint_prepared_prefix(task, outcome.commands, call.sequence, codec) descriptor = call.descriptor(codec) if call.recover: original = next((event.get("payload", {}) for event in reversed(history) @@ -1710,32 +1701,10 @@ async def _execute_prepared_local_activity( "prepared recovery is absent from this claim's canonical history", ) return refreshed - nonce = hashlib.sha256(json.dumps( - [task_id, task["run_id"], self.worker_id, epoch, call.sequence], - ).encode()).hexdigest() - started = time.monotonic() - receipt = await self.client.prepared_local_activity_operation( - task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, operation="prepare", - body={"sequence": call.sequence, "worker_attempt_id": nonce, "descriptor": descriptor}, - ) - attempt = PreparedAttempt.admitted( - receipt, task_id=task_id, run_id=task["run_id"], owner=self.worker_id, epoch=epoch, nonce=nonce, - heartbeat_timeout=call.command.heartbeat_timeout, cleanup=call.cleanup, request_started=started, - ) - handler = self.activities.get(call.command.activity_type) - if handler is None: - raise LocalActivityExecutionAborted("prepared local activity has no registered handler") - runner = PreparedLocalRunner( - self.client, attempt, shutdown=self._local_activity_shutdown, - observe=lambda value: task.update({"_prepared_cancellation_context": dict(value)}), - ) + runner, invocation = await self._prepare_local_callback(task, call, codec) + attempt = runner.attempt try: - receipt = await runner.execute(CallbackInvocation( - handler, tuple(call.command.arguments), ActivityInfo( - task_id, call.command.activity_type, attempt.attempt_id, attempt.attempt_number, - self.task_queue, self.worker_id, - ), {**task, "activity_attempt_id": attempt.attempt_id}, self.interceptors, - ), self._prepared_local_activity_processes) + receipt = await runner.execute(invocation, self._prepared_local_activity_processes) finally: if runner.callback is not None and runner.callback in self._prepared_local_activity_processes: self._abandoned_prepared_local_claims.add(task_id) @@ -1761,6 +1730,221 @@ async def _execute_prepared_local_activity( "prepared operation could not validate its original authority", ) from error + async def _checkpoint_prepared_prefix( + self, task: dict[str, Any], commands: list[Command], next_sequence: int, codec: str, + ) -> list[dict[str, Any]]: + wire = commands_to_server_commands(commands, self.task_queue, payload_codec=codec) + if any(command["type"] not in { + "record_side_effect", "record_version_marker", "upsert_memo", "upsert_search_attributes", + } for command in wire): + raise LocalActivityExecutionAborted("prepared prefix has no supported authored sequence range") + if any(command["type"] == "upsert_memo" for command in wire) and not self._workflow_memo_updates_supported: + raise LocalActivityExecutionAborted("prepared memo prefix requires negotiated workflow memo updates") + receipt = await self._checkpoint_prepared_commands(task, wire, next_sequence) + return await self._refresh_prepared_local_history(task, receipt) + + async def _checkpoint_prepared_commands( + self, task: dict[str, Any], commands: list[dict[str, Any]], next_sequence: int, *, group: bool = False, + ) -> dict[str, Any]: + task_id = task["task_id"] + epoch = task.get("workflow_task_attempt", 1) + start = next_sequence - len(commands) + if start < 1: + raise LocalActivityExecutionAborted("prepared checkpoint has no authored sequence range") + checkpoint = hashlib.sha256(json.dumps( + [task_id, self.worker_id, epoch, start, commands], sort_keys=True, allow_nan=False, + ).encode()).hexdigest() + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, + operation="checkpoint-group" if group else "checkpoint", + body={"checkpoint_id": checkpoint, "start_sequence": start, "commands": commands}, + ) + if (receipt.get("checkpointed") is not True or type(receipt.get("duplicate")) is not bool + or receipt.get("checkpoint_id") != checkpoint + or receipt.get("task_id") != task_id or receipt.get("workflow_run_id") != task["run_id"] + or type(receipt.get("workflow_task_attempt")) is not int + or receipt["workflow_task_attempt"] != epoch or receipt.get("lease_owner") != self.worker_id + or type(receipt.get("start_sequence")) is not int or receipt["start_sequence"] != start + or type(receipt.get("next_sequence")) is not int or receipt["next_sequence"] != next_sequence + or "reason" not in receipt or receipt["reason"] is not None): + raise LocalActivityExecutionAborted("prepared checkpoint lacks an original-claim receipt") + return receipt + + async def _prepare_local_callback( + self, task: dict[str, Any], call: PreparedLocalActivityCall, codec: str, + ) -> tuple[PreparedLocalRunner, CallbackInvocation]: + task_id = task["task_id"] + epoch = task.get("workflow_task_attempt", 1) + nonce = hashlib.sha256(json.dumps( + [task_id, task["run_id"], self.worker_id, epoch, call.sequence], + ).encode()).hexdigest() + started = time.monotonic() + receipt = await self.client.prepared_local_activity_operation( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=epoch, operation="prepare", + body={"sequence": call.sequence, "worker_attempt_id": nonce, "descriptor": call.descriptor(codec)}, + ) + attempt = PreparedAttempt.admitted( + receipt, task_id=task_id, run_id=task["run_id"], owner=self.worker_id, epoch=epoch, nonce=nonce, + heartbeat_timeout=call.command.heartbeat_timeout, cleanup=call.cleanup, request_started=started, + ) + handler = self.activities.get(call.command.activity_type) + if handler is None: + raise LocalActivityExecutionAborted("prepared local activity has no registered handler") + + def observe(value: Mapping[str, Any]) -> None: + previous = task.get("_prepared_cancellation_context") + if previous is not None and CancellationContext.from_dict(previous) != CancellationContext.from_dict(value): + raise LocalActivityExecutionAborted("prepared group members observed different cancellation authority") + task["_prepared_cancellation_context"] = dict(value) + + runner = PreparedLocalRunner(self.client, attempt, shutdown=self._local_activity_shutdown, observe=observe) + invocation = CallbackInvocation( + handler, tuple(call.command.arguments), ActivityInfo( + task_id, call.command.activity_type, attempt.attempt_id, attempt.attempt_number, + self.task_queue, self.worker_id, + ), {**task, "activity_attempt_id": attempt.attempt_id}, self.interceptors, + ) + return runner, invocation + + async def _execute_prepared_local_activity_group( + self, task: dict[str, Any], history: list[dict[str, Any]], outcome: ReplayOutcome, + ) -> list[dict[str, Any]]: + group = outcome.prepared_local_activity_group + if group is None or not self._prepared_local_activity_groups_supported: + raise LocalActivityExecutionAborted("prepared group lacks negotiated atomic admission") + codec = _validate_payload_codec(task.get("payload_codec")) or serializer.AVRO_CODEC + members: list[tuple[PreparedLocalRunner, CallbackInvocation]] = [] + executions: dict[asyncio.Task[dict[str, Any]], int] = {} + receipts: dict[int, dict[str, Any]] = {} + try: + if outcome.commands: + return await self._checkpoint_prepared_prefix(task, outcome.commands, group.base_sequence, codec) + if not group.committed: + calls = {call.sequence: call for call in group.calls} + commands: list[dict[str, Any]] = [] + for sequence, command in enumerate(group.commands, group.base_sequence): + if sequence in calls: + commands.append({**calls[sequence].descriptor(codec), "type": "prepare_local_activity"}) + else: + commands.extend(commands_to_server_commands([command], self.task_queue, payload_codec=codec)) + receipt = await self._checkpoint_prepared_commands( + task, commands, group.base_sequence + group.size, group=True, + ) + locals_ = receipt.get("local_activities") + identities: set[str] = set() + executions_by_sequence: dict[int, str] = {} + if not isinstance(locals_, list) or len(locals_) != len(group.calls): + raise LocalActivityExecutionAborted("atomic checkpoint omitted an authored local member") + for call, local in zip(group.calls, locals_, strict=True): + identity = local.get("activity_execution_id") if isinstance(local, Mapping) else None + if (not isinstance(local, Mapping) or type(local.get("sequence")) is not int + or local["sequence"] != call.sequence or not isinstance(identity, str) + or not identity.strip() or identity in identities): + raise LocalActivityExecutionAborted("atomic checkpoint changed its authored local members") + identities.add(identity) + executions_by_sequence[call.sequence] = identity + # Replay the entire canonical batch before preparing any callback. + refreshed = await self._refresh_prepared_local_history(task, receipt) + opening_kinds = { + "prepare_local_activity": {"ActivityScheduled"}, "schedule_activity": {"ActivityScheduled"}, + "start_timer": {"TimerScheduled"}, + "start_child_workflow": {"ChildWorkflowScheduled", "ChildRunStarted"}, + } + for sequence, wire_command in enumerate(commands, group.base_sequence): + opening = next((event for event in refreshed + if event.get("event_type", event.get("type")) in opening_kinds[wire_command["type"]] + and event.get("payload", {}).get("sequence", event.get("payload", {}).get( + "workflow_sequence")) == sequence), None) + payload = opening.get("payload", {}) if opening is not None else {} + if opening is None or payload.get("parallel_group_path") != wire_command["parallel_group_path"]: + raise LocalActivityExecutionAborted("atomic checkpoint is absent from canonical group history") + if sequence in calls and payload.get("activity_execution_id") != executions_by_sequence[sequence]: + raise LocalActivityExecutionAborted( + "atomic checkpoint changed canonical local execution identity", + ) + return refreshed + for call in group.calls: + if call.recover: + return await self._execute_prepared_local_activity( + task, history, ReplayOutcome([], prepared_local_activity=call), + ) + for call in group.calls: + members.append(await self._prepare_local_callback(task, call, codec)) + for index, (runner, invocation) in enumerate(members): + execution = asyncio.create_task(runner.execute(invocation, self._prepared_local_activity_processes)) + executions[execution] = index + pending = set(executions) + while pending: + done, pending = await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED) + for execution in done: + receipt = await execution + if receipt["claim_released"]: + raise _WorkflowClaimDeferred("Native released the group claim for a durable local retry") + receipts[executions[execution]] = receipt + # A cursor from a later-settled receipt includes all prior commits. + last = next(reversed(receipts.values())) + refreshed = await self._refresh_prepared_local_history(task, last) + for index, receipt in receipts.items(): + call = group.calls[index] + attempt = members[index][0].attempt + event = next((event for event in refreshed if event.get("id") == receipt["event_id"]), None) + payload = event.get("payload", {}) if event is not None else {} + if (event is None or event.get("event_type", event.get("type")) != receipt["event_type"] + or payload.get("sequence", payload.get("workflow_sequence")) != call.sequence + or payload.get("activity_execution_id") != attempt.execution_id + or payload.get("activity_attempt_id") != attempt.attempt_id): + raise LocalActivityExecutionAborted("prepared group outcome is absent from canonical history") + return refreshed + except PreparedCancellationObserved: + # Stop every sibling first, including admitted members that did not + # spawn. Only then report their canonical cancellation fences. + await self._join_prepared_group(executions, members, task) + for execution, index in executions.items(): + if not execution.cancelled() and execution.exception() is None: + receipts[index] = execution.result() + for index, (runner, _) in enumerate(members): + if index in receipts or runner.stop_acknowledged: + continue + try: + await runner.control() + except PreparedCancellationObserved: + await runner.acknowledge_stop() + continue + raise LocalActivityExecutionAborted( + "joined group member lacks its canonical cancellation fence", + ) from None + raise + except ServerError as error: + await self._join_prepared_group(executions, members, task) + if error.reason() == "cancellation_requested": + # No group process can start until all admissions have validated. + for runner, _ in members: + try: + await runner.control() + except PreparedCancellationObserved: + await runner.acknowledge_stop() + await self._renew_local_workflow_lease(task) + raise LocalActivityExecutionAborted("prepared group operation has an unknown or refused outcome") from error + except (_WorkflowClaimDeferred, LocalActivityExecutionAborted): + raise + except Exception as error: + raise LocalActivityExecutionAborted("prepared group could not validate original authority") from error + finally: + await self._join_prepared_group(executions, members, task) + + async def _join_prepared_group( + self, executions: Mapping[asyncio.Task[dict[str, Any]], int], + members: list[tuple[PreparedLocalRunner, CallbackInvocation]], task: dict[str, Any], + ) -> None: + for execution in executions: + if not execution.done(): + execution.cancel() + await asyncio.gather(*executions, return_exceptions=True) + if any(runner.callback is not None and runner.callback in self._prepared_local_activity_processes + for runner, _ in members): + self._abandoned_prepared_local_claims.add(task["task_id"]) + raise LocalActivityExecutionAborted("prepared group callback stop remains unconfirmed") + def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str, Any]: storage = self.external_storage threshold = self.external_storage_threshold_bytes diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 103a11e..78d0a02 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -572,6 +572,7 @@ class RecordLocalActivity: outcome: dict[str, Any] | None = field(default=None, init=False, repr=False) arguments_envelope: dict[str, Any] | None = field(default=None, init=False, repr=False) result_envelope: dict[str, Any] | None = field(default=None, init=False, repr=False) + _parallel_group_path: list[dict[str, Any]] | None = field(default=None, init=False, repr=False, compare=False) def __post_init__(self) -> None: if not isinstance(self.activity_type, str): @@ -658,9 +659,32 @@ def descriptor(self, payload_codec: str) -> dict[str, Any]: "request_id": self.cleanup["request_id"], "delivery_history_event_id": self.cleanup["delivery_history_event_id"], } + _apply_parallel_group_metadata(command, descriptor) return descriptor +@dataclass(frozen=True) +class PreparedLocalActivityGroup: + """A complete ordinary parallel group, with only unresolved local calls.""" + + base_sequence: int + size: int + commands: tuple[Command, ...] + calls: tuple[PreparedLocalActivityCall, ...] + committed: bool + + def __post_init__(self) -> None: + sequences = [call.sequence for call in self.calls] + if ( + self.base_sequence < 1 or not 1 <= self.size <= 100 or not self.calls + or len(self.commands) != (0 if self.committed else self.size) + or len(set(sequences)) != len(sequences) + or any(not self.base_sequence <= call.sequence < self.base_sequence + self.size + or not call.command._parallel_group_path for call in self.calls) + ): + raise LocalActivityExecutionAborted("prepared group lacks complete bounded authored membership") + + class LocalActivityExecutionAborted(Exception): """The workflow task lease could not be trusted after local execution began.""" @@ -2313,6 +2337,7 @@ class ReplayOutcome: message_stream_waits: list[dict[str, Any]] = field(default_factory=list) cancellation_delivery: CancellationDelivery | None = None prepared_local_activity: PreparedLocalActivityCall | None = None + prepared_local_activity_group: PreparedLocalActivityGroup | None = None class Replayer: @@ -2687,6 +2712,7 @@ def replay( cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, ) -> ReplayOutcome: return _replay_state( workflow_cls, @@ -2702,6 +2728,7 @@ def replay( cancellation_request=cancellation_request, local_activity_executor=local_activity_executor, prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, ).outcome @@ -3702,8 +3729,11 @@ def _replay_state( cancellation_request: Mapping[str, Any] | None = None, local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, stop_at_uncommitted_cancellation: bool = False, ) -> _ReplayState: + if prepare_local_activity_groups and not prepare_local_activities: + raise LocalActivityExecutionAborted("prepared groups require prepared local activity admission") if payload_codec is not None and payload_codec != serializer.AVRO_CODEC: try: serializer.decode("", codec=payload_codec) @@ -3847,12 +3877,14 @@ def _replay_state( def _state( commands: list[Command], prepared_local_activity: PreparedLocalActivityCall | None = None, + prepared_local_activity_group: PreparedLocalActivityGroup | None = None, ) -> _ReplayState: return _ReplayState( outcome=ReplayOutcome( commands=commands, cancellation_delivery=cancellation_intent, prepared_local_activity=prepared_local_activity, + prepared_local_activity_group=prepared_local_activity_group, message_stream_cursors=[ {"stream_name": name, "through_position": position} for name, position in sorted(ctx._message_stream_cursors.items()) @@ -4719,7 +4751,7 @@ def _same_logical_condition_wait(opened: Mapping[str, Any], cmd: WaitCondition) terminal_condition_reopen_cmd: WaitCondition | None = None def _parallel_leaf_kind(command: Any) -> str: - if isinstance(command, ScheduleActivity): + if isinstance(command, ScheduleActivity | RecordLocalActivity): return "activity" if isinstance(command, StartTimer): return "timer" @@ -5269,6 +5301,91 @@ def _contains_local_activity(operation: Any) -> bool: return any(_contains_local_activity(member) for _, member in operation.operations) return False + def _prepared_call(command: RecordLocalActivity, sequence: int) -> PreparedLocalActivityCall: + cleanup: dict[str, str] | None = None + if cancellation.request is not None: + delivery = cancellation.delivery + context = cancellation.request.context + delivery_id = ( + events[cancellation.delivery_index].get("id") + if cancellation.delivery_index is not None else None + ) + if ( + not cancellation_consumed or ctx._cancellation_shield_depth < 1 + or delivery is None or context is None + or sequence < delivery.sequence + delivery.sequence_span + or not isinstance(delivery_id, str) or not delivery_id.strip() + ): + raise LocalActivityExecutionAborted( + "prepared cleanup requires a shield after canonical cancellation delivery", + ) + cleanup = { + "request_id": context.request_id, "root_request_id": context.root_request_id, + "delivery_history_event_id": delivery_id, + "cleanup_deadline_at": context.to_dict()["cleanup_deadline_at"], + } + started = False + for event in events: + if _workflow_sequence(event.get("payload") or {}) != sequence: + continue + if _history_event_type(event) == "ActivityStarted": + started = True + elif _history_event_type(event) == "ActivityRetryScheduled": + started = False + return PreparedLocalActivityCall(command, sequence, started, cleanup) + + def _prepared_group(commands: list[Any]) -> PreparedLocalActivityGroup | None: + if not 1 <= len(commands) <= 100: + raise LocalActivityExecutionAborted("atomic prepared groups require 1 to 100 authored members") + openings = { + "ActivityScheduled": "activity", "TimerScheduled": "timer", + "ChildWorkflowScheduled": "child workflow", "ChildRunStarted": "child workflow", + } + present: set[int] = set() + for offset, command in enumerate(commands): + sequence = current_call_sequence + offset + if not isinstance(command, ScheduleActivity | RecordLocalActivity | StartTimer | StartChildWorkflow): + raise LocalActivityExecutionAborted("prepared group contains an unsupported admission operation") + for event in events: + payload = event.get("payload") or {} + if _workflow_sequence(payload) != sequence: + continue + kind = _history_event_type(event) + if kind not in openings: + continue + # A complete path is admission authority. Legacy metadata-poor + # history cannot authorize fresh Source callbacks. + details = _recorded_step_details(payload) + if payload.get("parallel_group_path") != command._parallel_group_path: + raise NonDeterministicReplayError( + sequence, "complete authored parallel group path", [kind], + detail="prepared group history changed or omitted its membership", + ) + _assert_step_matches(command, _RecordedStep(sequence, openings[kind], [kind], details)) + present.add(sequence) + if present and len(present) != len(commands): + raise NonDeterministicReplayError( + current_call_sequence, "all parallel members scheduled", [], + detail="prepared group history is missing a declared member", + ) + # Never reinterpret terminal/started history without its atomic opening. + group_sequences = set(range(current_call_sequence, current_call_sequence + len(commands))) + if not present and group_sequences.intersection(event_types_by_sequence): + raise NonDeterministicReplayError( + current_call_sequence, "complete atomic group opening", [], + detail="prepared group history lacks its original scheduling batch", + ) + calls = tuple( + _prepared_call(command, current_call_sequence + offset) + for offset, command in enumerate(commands) + if isinstance(command, RecordLocalActivity) and current_call_sequence + offset not in resolved_sequences + ) + if not calls: + return None + return PreparedLocalActivityGroup( + current_call_sequence, len(commands), () if present else tuple(commands), calls, bool(present), + ) + def _assert_cancellation_call_matches(command: Any, boundary: CancellationDelivery) -> None: if isinstance(command, list): leaves, _ = _annotate_parallel_commands(command, boundary.sequence) @@ -5362,7 +5479,9 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl return _terminal_state(stop.value, include_pending=True) first = False _apply_due_receivers() - if prepare_local_activities and isinstance(cmd, list | SelectGroup) and _contains_local_activity(cmd): + if prepare_local_activities and isinstance(cmd, list | SelectGroup) and _contains_local_activity(cmd) and ( + not prepare_local_activity_groups or isinstance(cmd, SelectGroup) + ): raise LocalActivityExecutionAborted( "prepared local groups require an implemented atomic group consumer", ) @@ -5510,8 +5629,22 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl "ordinary parallel list groups do not support WaitCondition; " "use WorkflowContext.select() for durable condition selection" ) - cmd, result_shape = _annotate_parallel_commands(cmd) + prepared_group = prepare_local_activities and _contains_local_activity(cmd) + cmd, result_shape = _annotate_parallel_commands(cmd, current_call_sequence if prepared_group else None) needed = len(cmd) + if prepared_group: + group = _prepared_group(cmd) + if group is not None: + ctx.logger._set_replaying(False) + return _state(pending, prepared_local_activity_group=group) + if any(current_call_sequence + offset not in resolved_sequences for offset in range(needed)): + # Local results are durable. Remote members own their + # remaining work, with no repeat scheduling batch. + return _state(pending) + elif any(isinstance(command, RecordLocalActivity) for command in cmd) and ( + result_cursor + needed > len(resolved_results) + ): + raise LocalActivityExecutionAborted("local groups require negotiated atomic prepared admission") if result_cursor + needed <= len(resolved_results): for offset, child_command in enumerate(cmd): _assert_next_step_matches(child_command, offset) @@ -5819,37 +5952,7 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl ctx.logger._set_replaying(False) _assert_pending_step_matches(cmd) if isinstance(cmd, RecordLocalActivity) and prepare_local_activities: - cleanup: dict[str, str] | None = None - if cancellation.request is not None: - delivery = cancellation.delivery - context = cancellation.request.context - delivery_id = ( - events[cancellation.delivery_index].get("id") - if cancellation.delivery_index is not None else None - ) - if ( - not cancellation_consumed or ctx._cancellation_shield_depth < 1 - or delivery is None or context is None - or current_call_sequence < delivery.sequence + delivery.sequence_span - or not isinstance(delivery_id, str) or not delivery_id.strip() - ): - raise LocalActivityExecutionAborted( - "prepared cleanup requires a shield after canonical cancellation delivery", - ) - cleanup = { - "request_id": context.request_id, "root_request_id": context.root_request_id, - "delivery_history_event_id": delivery_id, - "cleanup_deadline_at": context.to_dict()["cleanup_deadline_at"], - } - started = False - for event in events: - if _workflow_sequence(event.get("payload") or {}) != current_call_sequence: - continue - if _history_event_type(event) == "ActivityStarted": - started = True - elif _history_event_type(event) == "ActivityRetryScheduled": - started = False - return _state(pending, PreparedLocalActivityCall(cmd, current_call_sequence, started, cleanup)) + return _state(pending, _prepared_call(cmd, current_call_sequence)) if isinstance(cmd, RecordLocalActivity) and local_activity_executor is not None: local_result = local_activity_executor(cmd) if cmd.outcome is None or cmd.arguments_envelope is None: diff --git a/tests/fixtures/replay_regressions/prepared-local-group-cold-results.json b/tests/fixtures/replay_regressions/prepared-local-group-cold-results.json new file mode 100644 index 0000000..f160ba6 --- /dev/null +++ b/tests/fixtures/replay_regressions/prepared-local-group-cold-results.json @@ -0,0 +1,23 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "prepared-local-group-cold-results", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "tests.replay.prepared-local-group-cold-results", + "input": [], + "payload_codec": "avro" + }, + "history": [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 0}]}}, + {"event_type": "ActivityScheduled", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 1}]}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "activity_execution_id": "first", "activity_attempt_id": "first-original-attempt", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 0}]}}, + {"event_type": "ActivityStarted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "activity_execution_id": "second", "activity_attempt_id": "second-original-attempt", "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 1}]}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 1, "activity_type": "prepared.first", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}, "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 0}]}}, + {"event_type": "ActivityCompleted", "payload": {"sequence": 2, "activity_type": "prepared.second", "execution_mode": "local", "result": {"codec": "avro", "blob": "wwHioz3/VYAiNwoUZmFzdC12YWx1ZQ=="}, "parallel_group_path": [{"parallel_group_id": "parallel-activities:1:2", "parallel_group_kind": "activity", "parallel_group_base_sequence": 1, "parallel_group_size": 2, "parallel_group_index": 1}]}} + ], + "expected": { + "command_sequence": [{"type": "complete_workflow", "command_type": "CompleteWorkflow", "result": ["fast-value", "fast-value"]}] + } +} diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index 9ac0941..b54b4d3 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -28,6 +28,7 @@ async def prepared_runtime(server_url: str, server_token: str, monkeypatch: pyte async with Client(server_url, token=server_token, namespace="default") as client: info = await client.get_cluster_info() assert info["worker_protocol"]["server_capabilities"]["prepared_local_activities"] is True + assert info["worker_protocol"]["server_capabilities"]["prepared_local_activity_groups"] is True assert info["worker_protocol"]["version"] == "1.20" monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") @@ -53,15 +54,25 @@ async def prepared_local(marker: str, phase: str) -> dict[str, Any]: @workflow.defn(name="tests.python-prepared-cancellation") class PreparedCancellationWorkflow: - def run(self, ctx: Any, marker: str) -> Any: + def run(self, ctx: Any, marker: str, group: bool = False) -> Any: try: - yield ctx.local_activity("tests.python-prepared-blocked", [marker, "work", None]) + if group: + yield [ctx.local_activity("tests.python-prepared-blocked", [marker, "work-0", None]), + ctx.local_activity("tests.python-prepared-blocked", [marker, "work-1", None])] + else: + yield ctx.local_activity("tests.python-prepared-blocked", [marker, "work", None]) except WorkflowCancelled as error: with ctx.cancellation_shield(): - yield ctx.local_activity( - "tests.python-prepared-blocked", [marker, "cleanup", error.request_id], - retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, - ) + if group: + yield [ctx.local_activity( + "tests.python-prepared-blocked", [marker, "cleanup-" + str(index), error.request_id], + retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, + ) for index in range(2)] + else: + yield ctx.local_activity( + "tests.python-prepared-blocked", [marker, "cleanup", error.request_id], + retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, + ) return error.request_id return "not cancelled" @@ -74,18 +85,39 @@ async def prepared_blocked(marker: str, phase: str, request_id: str | None) -> d pending.write_text(json.dumps({"phase": phase, "callback_pid": os.getpid(), "request_id": request_id, "activity_attempt_id": info.activity_attempt_id})) pending.replace(path) - if phase == "work" or os.environ.get("DW_PREPARED_FIXTURE_MODE") == "hold": + if phase.startswith("work") or os.environ.get("DW_PREPARED_FIXTURE_MODE") == "hold": # Intentionally no application heartbeat. await asyncio.Event().wait() return {"request_id": request_id, "bytes": b"\x00\xff"} +@workflow.defn(name="tests.python-prepared-group") +class PreparedGroupWorkflow: + def run(self, ctx: Any, marker: str) -> Any: + yield ctx.upsert_memo({"before": "atomic local group"}) + first = ctx.local_activity("tests.python-prepared-peer", [marker, "first", "second"], heartbeat_timeout=10) + second = ctx.local_activity("tests.python-prepared-peer", [marker, "second", "first"], heartbeat_timeout=10) + return (yield [first, [second, ctx.start_timer(1)]]) + + +@activity.defn(name="tests.python-prepared-peer") +async def prepared_peer(marker: str, phase: str, peer: str) -> dict[str, Any]: + info = activity.context().info + Path(marker + "." + phase).write_text(json.dumps({ + "callback_pid": os.getpid(), "attempt_id": info.activity_attempt_id, + })) + await remote_marker(Path(marker + "." + peer)) + await activity.context().heartbeat({"phase": phase}) + return {"phase": phase, "bytes": b"\x00\xff", "attempt": info.activity_attempt_id} + + def prepared_worker(client: Client, queue: str, **kwargs: Any) -> Worker: return Worker( client, task_queue=queue, worker_id=kwargs.pop("worker_id", queue + "-owner"), - workflows=[PreparedSequentialWorkflow, PreparedCancellationWorkflow], - activities=[prepared_local, prepared_blocked], - capabilities=["cooperative_cancellation", "prepared_local_activities"], **kwargs, + workflows=[PreparedSequentialWorkflow, PreparedCancellationWorkflow, PreparedGroupWorkflow], + activities=[prepared_local, prepared_blocked, prepared_peer], + capabilities=["cooperative_cancellation", "prepared_local_activities", "prepared_local_activity_groups"], + **kwargs, ) @@ -136,8 +168,43 @@ async def test_prepared_prefix_two_callbacks_and_cold_replay_use_canonical_histo await worker.stop() -async def test_prepared_callback_stop_cleanup_sigkill_and_cold_recovery_keep_original_30_second_deadline( +async def test_prepared_nested_mixed_group_starts_peers_concurrently_and_cold_replays_results_in_position( server_url: str, server_token: str, tmp_path: Path, +) -> None: + queue = "py-prepared-group-" + uuid.uuid4().hex[:8] + marker = str(tmp_path / "callback") + async with Client(server_url, token=server_token, namespace="default") as client: + worker = prepared_worker(client, queue) + owner = asyncio.create_task(worker.run()) + try: + handle = await client.start_workflow(workflow_type="tests.python-prepared-group", task_queue=queue, + workflow_id=queue, input=[marker]) + result = await handle.result(timeout=15) + assert result[0]["phase"] == "first" and result[1][0]["phase"] == "second" + assert result[0]["bytes"] == result[1][0]["bytes"] == b"\x00\xff" + assert result[0]["attempt"] != result[1][0]["attempt"] + history = await events(handle) + kinds = [event["event_type"] for event in history] + assert kinds.count("ActivityScheduled") == kinds.count("ActivityStarted") == 2 + assert kinds.count("ActivityCompleted") == 2 + assert kinds.count("ActivityHeartbeatRecorded") == 2 + assert kinds.count("TimerScheduled") == kinds.count("TimerFired") == 1 + assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 + completed = replay(PreparedGroupWorkflow, history, [marker], run_id=handle.run_id or "", + prepare_local_activities=True, prepare_local_activity_groups=True) + assert completed.prepared_local_activity_group is None + assert completed.commands[0].result == result # type: ignore[union-attr] + for phase in ("first", "second"): + await callback_gone(json.loads(Path(marker + "." + phase).read_text())["callback_pid"]) + print("prepared nested mixed group history: " + json.dumps(history)) + finally: + await worker.stop() + await asyncio.wait_for(owner, timeout=10) + + +@pytest.mark.parametrize("group", [False, True], ids=["sequential", "atomic-group"]) +async def test_prepared_callback_stop_cleanup_sigkill_and_cold_recovery_keep_original_30_second_deadline( + server_url: str, server_token: str, tmp_path: Path, group: bool, ) -> None: queue = "py-prepared-cancel-" + uuid.uuid4().hex[:8] marker = str(tmp_path / "callback") @@ -158,20 +225,26 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: async with Client(server_url, token=server_token, namespace="default") as client: try: handle = await client.start_workflow(workflow_type="tests.python-prepared-cancellation", task_queue=queue, - workflow_id=queue, input=[marker]) + workflow_id=queue, input=[marker, group]) first = await owner(queue + "-first", "hold") - work = await remote_marker(Path(marker + ".work." + queue + "-first")) + work_phases = ["work-0", "work-1"] if group else ["work"] + cleanup_phases = ["cleanup-0", "cleanup-1"] if group else ["cleanup"] + work = [await remote_marker(Path(marker + "." + phase + "." + queue + "-first")) + for phase in work_phases] accepted = await handle.request_cancellation(cleanup_timeout_seconds=30) original = accepted["cancellation_request"] requested_history = await events(handle) original_context = next(event["payload"]["cancellation"] for event in requested_history if event["event_type"] == "CooperativeCancellationRequested") - cleanup = await remote_marker(Path(marker + ".cleanup." + queue + "-first")) - assert cleanup["request_id"] == original["request_id"] - await callback_gone(work["callback_pid"]) + cleanup = [await remote_marker(Path(marker + "." + phase + "." + queue + "-first")) + for phase in cleanup_phases] + assert all(item["request_id"] == original["request_id"] for item in cleanup) + for item in work: + await callback_gone(item["callback_pid"]) first.kill() await asyncio.wait_for(first.wait(), timeout=10) - await callback_gone(cleanup["callback_pid"]) + for item in cleanup: + await callback_gone(item["callback_pid"]) duplicate = await handle.request_cancellation(cleanup_timeout_seconds=90) for field in ("request_id", "requested_at", "cleanup_deadline_at"): assert duplicate["cancellation_request"][field] == original[field] @@ -184,24 +257,31 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: history = await events(handle) kinds = [event["event_type"] for event in history] assert kinds.count("CooperativeCancellationDelivered") == kinds.count("WorkflowCancelled") == 1 - assert kinds.count("ActivityCancellationAcknowledged") == 1 + assert kinds.count("ActivityCancellationAcknowledged") == len(work) + assert kinds.count("ActivityCancelled") == len(work) assert "ActivityHeartbeatRecorded" not in kinds terminal = next(event for event in history if event["event_type"] == "WorkflowCancelled") assert datetime.fromisoformat(terminal["timestamp"].replace("Z", "+00:00")) < deadline receipts = [json.loads(line) for line in trace_path.read_text().splitlines()] admissions = [entry["receipt"] for entry in receipts if entry["operation"] == "prepare" and entry["receipt"].get("cancellation_cleanup")] - assert len(admissions) == 2 - assert admissions[0]["activity_attempt_id"] != admissions[1]["activity_attempt_id"] - assert admissions[0]["cancellation_cleanup"] == admissions[1]["cancellation_cleanup"] + assert len(admissions) == 2 * len(cleanup) + assert len({item["activity_attempt_id"] for item in admissions}) == len(admissions) + assert all(item["cancellation_cleanup"] == admissions[0]["cancellation_cleanup"] for item in admissions) authority = admissions[0]["cancellation_cleanup"] assert authority["request_id"] == original["request_id"] assert authority["root_request_id"] == original_context["root_request_id"] assert datetime.fromisoformat(authority["cleanup_deadline_at"].replace("Z", "+00:00")) == deadline recoveries = [entry["receipt"] for entry in receipts if entry["operation"] == "recover"] - assert len(recoveries) == 1 and recoveries[0]["callback_stop_state"] == "unknown" - replacement = await remote_marker(Path(marker + ".cleanup." + queue + "-replacement")) - await callback_gone(replacement["callback_pid"]) + assert len(recoveries) == len(cleanup) + assert all(item["callback_stop_state"] == "unknown" for item in recoveries) + replacement = [await remote_marker(Path(marker + "." + phase + "." + queue + "-replacement")) + for phase in cleanup_phases] + for item in replacement: + await callback_gone(item["callback_pid"]) + print("prepared physically joined callbacks: " + json.dumps({ + "work": work, "killed_cleanup": cleanup, "replacement_cleanup": replacement, + })) print("prepared cleanup recovery receipts: " + json.dumps(receipts)) print("prepared cleanup recovery history: " + json.dumps(history)) finally: diff --git a/tests/test_activity_process.py b/tests/test_activity_process.py index 82e3d14..014b7f4 100644 --- a/tests/test_activity_process.py +++ b/tests/test_activity_process.py @@ -227,3 +227,37 @@ def callback() -> None: pass with pytest.raises(ValueError, match="spawn-compatible.*importable handler"): SupervisedCallback(invocation(callback)) + + +async def test_cancellation_during_confirmed_result_join_still_reaps_and_proves_stop( + monkeypatch: pytest.MonkeyPatch, +) -> None: + callback = SupervisedCallback(invocation(typed_result, b"typed")) + entering_join = asyncio.Event() + release_join = asyncio.Event() + close = SupervisedCallback.close + + async def paused_close(self: SupervisedCallback) -> None: + entering_join.set() + await release_join.wait() + await close(self) + + monkeypatch.setattr(SupervisedCallback, "close", paused_close) + result: asyncio.Task[Any] | None = None + try: + await callback.start() + result = asyncio.create_task(callback.result(no_heartbeat)) + await asyncio.wait_for(entering_join.wait(), timeout=5) + result.cancel() + await asyncio.sleep(0) + assert not callback.stopped + release_join.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(result, timeout=7) + assert callback.stopped and callback._joined and callback._exitcode == 0 + finally: + release_join.set() + if result is not None: + await asyncio.gather(result, return_exceptions=True) + if not callback.stopped: + await callback.stop() diff --git a/tests/test_prepared_local_activity_groups.py b/tests/test_prepared_local_activity_groups.py new file mode 100644 index 0000000..3304295 --- /dev/null +++ b/tests/test_prepared_local_activity_groups.py @@ -0,0 +1,390 @@ +from __future__ import annotations + +import asyncio +import json +import os +from copy import deepcopy +from pathlib import Path +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import activity, serializer, workflow +from durable_workflow._prepared_local_activity import PreparedCancellationObserved +from durable_workflow.client import Client +from durable_workflow.errors import NonDeterministicReplayError, ServerError, WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted, replay +from tests.test_activity_process import blocked_without_python_progress, wait_for_exit, wait_for_file +from tests.test_cancellation_context import delivery, request, snapshot +from tests.test_prepared_local_activity_worker import admission, control, timestamp +from tests.test_replay_regression_corpus import PreparedLocalGroupColdResultsWorkflow +from tests.test_worker import compatible_cluster_info + + +@workflow.defn(name="prepared.two-locals") +class TwoLocals: + def run(self, ctx: workflow.WorkflowContext, marker: str, prefix: bool = False): # type: ignore[no-untyped-def] + if prefix: + yield ctx.upsert_memo({"before": "group"}) + return (yield [ctx.local_activity("prepared.group", [marker + ".0", marker + ".1"]), + ctx.local_activity("prepared.group", [marker + ".1", marker + ".0"])]) + + +@workflow.defn(name="prepared.mixed-group") +class MixedGroup: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + return (yield [ctx.start_child_workflow("python.child", [], task_queue="child"), + [ctx.schedule_activity("remote", []), ctx.local_activity("prepared.group", [])]]) + + +@workflow.defn(name="prepared.group-cleanup") +class GroupCleanup: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled: + with ctx.cancellation_shield(): + yield [ctx.local_activity("prepared.group", []), ctx.local_activity("prepared.group", [])] + return "cleaned" + + +@activity.defn(name="prepared.group") +async def concurrent_callback(marker: str, peer: str) -> dict[str, Any]: + Path(marker).write_text(str(os.getpid())) + for _ in range(200): + if Path(peer).exists(): + return {"bytes": b"\x00\xff", "attempt": activity.context().info.activity_attempt_id} + await asyncio.sleep(0.01) + raise RuntimeError("parallel sibling did not start") + + +def capture(cls: type = TwoLocals, history: list[dict[str, Any]] | None = None, inputs: list[Any] | None = None): # type: ignore[no-untyped-def] + return replay(cls, history or [], ["unused"] if inputs is None else inputs, + prepare_local_activities=True, prepare_local_activity_groups=True, run_id="run-1", + local_activity_executor=lambda _: pytest.fail("callback ran without group admission")) + + +def local_event(kind: str, sequence: int, *, base: int = 1, **extra: Any) -> dict[str, Any]: + metadata = {"parallel_group_id": f"parallel-activities:{base}:2", "parallel_group_kind": "activity", + "parallel_group_base_sequence": base, "parallel_group_size": 2, + "parallel_group_index": sequence - base} + return {"id": kind + str(sequence), "event_type": kind, "payload": { + "sequence": sequence, "activity_type": "prepared.group", "execution_mode": "local", + "activity_execution_id": f"execution-{sequence}", "activity_attempt_id": f"attempt-{sequence}", + **metadata, "parallel_group_path": [metadata], **extra, + }} + + +def test_complete_nested_mixed_batch_is_captured_without_inventing_local_outcomes() -> None: + group = capture(MixedGroup, inputs=[]).prepared_local_activity_group + assert group is not None and not group.committed and group.base_sequence == 1 and group.size == 3 + assert [call.sequence for call in group.calls] == [3] + descriptor = group.calls[0].descriptor("avro") + assert "outcome" not in descriptor and "result" not in descriptor and "queue" not in descriptor + assert len(descriptor["parallel_group_path"]) == 2 + assert descriptor["parallel_group_path"][0]["parallel_group_index"] == 2 + assert descriptor["parallel_group_path"][1]["parallel_group_index"] == 1 + assert group.commands[0].task_queue == "child" # type: ignore[union-attr] + + +def test_group_prefix_is_checkpointed_separately_and_preserves_authored_sequences() -> None: + outcome = capture(inputs=["unused", True]) + assert [type(command).__name__ for command in outcome.commands] == ["UpsertMemo"] + assert outcome.prepared_local_activity_group is not None + assert outcome.prepared_local_activity_group.base_sequence == 2 + assert [call.sequence for call in outcome.prepared_local_activity_group.calls] == [2, 3] + + +@pytest.mark.parametrize("history", [ + [local_event("ActivityScheduled", 1)], + [local_event("ActivityStarted", 1), local_event("ActivityStarted", 2)], + [{**local_event("ActivityScheduled", 1), "payload": { + **local_event("ActivityScheduled", 1)["payload"], "parallel_group_path": []}}, + local_event("ActivityScheduled", 2)], + [local_event("ActivityScheduled", 1, base=2), local_event("ActivityScheduled", 2, base=2)], +]) +def test_partial_or_changed_group_history_cannot_authorize_any_callback(history: list[dict[str, Any]]) -> None: + with pytest.raises(NonDeterministicReplayError): + capture(history=history) + + +def test_partial_cold_replay_skips_a_completed_member_and_recovers_only_the_started_sibling() -> None: + history = [local_event("ActivityScheduled", 1), local_event("ActivityScheduled", 2), + local_event("ActivityCompleted", 1, result=serializer.envelope("first")), + local_event("ActivityStarted", 2)] + group = capture(history=history).prepared_local_activity_group + assert group is not None and group.committed and group.commands == () + assert [call.sequence for call in group.calls] == [2] and group.calls[0].recover is True + history.append(local_event("ActivityRetryScheduled", 2)) + assert capture(history=history).prepared_local_activity_group.calls[0].recover is False + history.append(local_event("ActivityCompleted", 2, result=serializer.envelope("second"))) + outcome = capture(history=history) + assert outcome.prepared_local_activity_group is None and outcome.commands[0].result == ["first", "second"] + + +def test_reverse_completion_order_replays_results_at_the_authored_positions() -> None: + history = [local_event("ActivityScheduled", 1), local_event("ActivityScheduled", 2), + local_event("ActivityCompleted", 2, result=serializer.envelope("second")), + local_event("ActivityCompleted", 1, result=serializer.envelope("first"))] + assert capture(history=history).commands[0].result == ["first", "second"] + + +def test_cleanup_group_preserves_original_root_delivery_and_deadline_for_every_member() -> None: + history = [request(), {**delivery(), "id": "canonical-delivery"}] + group = capture(GroupCleanup, history, []).prepared_local_activity_group + assert group is not None and group.base_sequence == 2 + for call in group.calls: + assert call.cleanup == {"request_id": "request-1", "root_request_id": "root-1", + "delivery_history_event_id": "canonical-delivery", + "cleanup_deadline_at": snapshot()["cleanup_deadline_at"]} + assert group.calls[0].cleanup == group.calls[1].cleanup + + +@pytest.mark.parametrize("selection", [False, True]) +def test_unsupported_group_shapes_are_refused_before_callbacks(selection: bool) -> None: + @workflow.defn(name="prepared.unsupported") + class Unsupported: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + local = ctx.local_activity("prepared.group", []) + yield ctx.select([local]) if selection else [local] * 101 + + with pytest.raises(LocalActivityExecutionAborted): + capture(Unsupported, inputs=[]) + + +class GroupServer: + def __init__(self, markers: tuple[Path, Path]) -> None: + self.client = AsyncMock(spec=Client) + self.history: list[dict[str, Any]] = [] + self.trace: list[tuple[str, Any]] = [] + self.attempts: dict[str, dict[str, Any]] = {} + self.markers = markers + self.stop = False + self.bad_checkpoint = False + self.bad_second_admission = False + self.partial_checkpoint_history = False + self.lost_outcome = False + self.admission_cancelled = False + self.omit_second_outcome = False + self.client.get_cluster_info.return_value = compatible_cluster_info(worker_protocol={ + "version": "1.20", "server_capabilities": { + "prepared_local_activities": True, "prepared_local_activity_groups": True, + "cooperative_cancellation": True, + }, + }) + self.client.register_worker.return_value = {"registered": True} + self.client.complete_workflow_task.return_value = {"outcome": "completed"} + self.client.prepared_local_activity_operation.side_effect = self.operation + self.client.workflow_task_history.side_effect = self.page + + async def page(self, **_: Any) -> dict[str, Any]: + self.trace.append(("history", None)) + return {"history_events": deepcopy(self.history), "next_history_page_token": None} + + async def operation(self, **kwargs: Any) -> dict[str, Any]: + assert kwargs["task_id"] == "task" and kwargs["lease_owner"] == "owner" + assert kwargs["workflow_task_attempt"] == 3 + name = kwargs["operation"] + body = kwargs["body"] + attempt_id = kwargs.get("activity_attempt_id") + self.trace.append((name, body.get("sequence", attempt_id))) + if name == "checkpoint-group": + assert all(not path.exists() for path in self.markers) + assert [command["type"] for command in body["commands"]] == ["prepare_local_activity"] * 2 + self.history = [local_event("ActivityScheduled", 1), local_event("ActivityScheduled", 2)] + if self.partial_checkpoint_history: + self.history.pop() + locals_ = [{"sequence": index, "activity_execution_id": f"execution-{index}"} for index in (1, 2)] + if self.bad_checkpoint: + locals_[1]["activity_execution_id"] = "execution-1" + return {"checkpointed": True, "duplicate": False, "reason": None, + "checkpoint_id": body["checkpoint_id"], "task_id": "task", "workflow_run_id": "run-1", + "workflow_task_attempt": 3, "lease_owner": "owner", "start_sequence": 1, "next_sequence": 3, + "local_activities": locals_, "history_refresh_page_token": "canonical-page"} + if name == "prepare": + assert all(not path.exists() for path in self.markers) + sequence = body["sequence"] + if sequence == 2 and self.admission_cancelled: + self.stop = True + raise ServerError(409, {"reason": "cancellation_requested"}) + receipt = admission( + activity_execution_id=f"execution-{sequence}", activity_attempt_id=f"attempt-{sequence}", + worker_attempt_id=body["worker_attempt_id"], + ) + if sequence == 2 and self.bad_second_admission: + receipt["activity_attempt_id"] = "" + self.attempts[f"attempt-{sequence}"] = receipt + self.history.append(local_event("ActivityStarted", sequence)) + return receipt + assert isinstance(attempt_id, str) and attempt_id in self.attempts + receipt = self.attempts[attempt_id] + sequence = int(attempt_id[-1]) + if name == "control": + assert body == {"renew_lease": True} + if not self.stop: + return control(receipt) + context = {**snapshot(), "requested_at": timestamp(), "cleanup_deadline_at": timestamp(30)} + # One immutable context for the entire group. + if not hasattr(self, "context"): + self.context = context + return control(receipt, active=False, renewed=False, stop_required=True, reason="cancellation_requested", + fenced=True, cancellation_request=self.context, cancellation_history_event_id="cancelled", + history_refresh_page_token="canonical-page") + if name == "acknowledge-cancellation": + assert body == {"request_id": "request-1"} + for marker in self.markers: + # This member must have physically exited before its own ACK. + if marker.exists() and marker == self.markers[sequence - 1]: + with pytest.raises(ProcessLookupError): + os.kill(int(marker.read_text()), 0) + return {"acknowledged": True, "duplicate": False, "reason": None, + "history_event_id": "joined-" + attempt_id} + assert name == "outcome" + assert len(self.attempts) == 2 + if self.lost_outcome: + raise TimeoutError("canonical acknowledgment lost") + if sequence != 2 or not self.omit_second_outcome: + self.history.append(local_event("ActivityCompleted", sequence, result=body["report"]["result"])) + return {**receipt, "recorded": True, "event_id": "ActivityCompleted" + str(sequence), + "event_type": "ActivityCompleted", "workflow_run_id": "run-1", "recorded_at": timestamp(), + "claim_released": False, "created_task_ids": [], "history_refresh_page_token": "canonical-page"} + + async def worker(self, monkeypatch: pytest.MonkeyPatch, *, handler: Any = concurrent_callback) -> Worker: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + worker = Worker(self.client, task_queue="queue", worker_id="owner", workflows=[TwoLocals], + capabilities=["cooperative_cancellation", "prepared_local_activities", + "prepared_local_activity_groups"]) + await worker._register() + worker.activities["prepared.group"] = handler + return worker + + def task(self) -> dict[str, Any]: + return {"task_id": "task", "workflow_type": "prepared.two-locals", "workflow_id": "workflow-1", + "run_id": "run-1", "workflow_task_attempt": 3, "payload_codec": "avro", + "arguments": serializer.envelope([str(self.markers[0])[:-2]]), "history_events": deepcopy(self.history)} + + +def server(tmp_path: Path) -> GroupServer: + return GroupServer((tmp_path / "callback.0", tmp_path / "callback.1")) + + +async def test_atomic_batch_then_complete_admission_precedes_concurrent_callbacks_and_canonical_results( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch) + commands = await worker._run_workflow_task(fixture.task()) + assert commands is not None and [command["type"] for command in commands] == ["complete_workflow"] + result = serializer.decode_envelope(commands[0]["result"]) + assert [item["attempt"] for item in result] == ["attempt-1", "attempt-2"] + assert all(item["bytes"] == b"\x00\xff" for item in result) + assert fixture.trace[:4] == [("checkpoint-group", None), ("history", None), ("prepare", 1), ("prepare", 2)] + for marker in fixture.markers: + await wait_for_exit(int(marker.read_text())) + assert not worker._prepared_local_activity_processes + fixture.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("failure", ["bad_checkpoint", "partial_checkpoint_history", "bad_second_admission"]) +async def test_incomplete_admission_never_spawns_or_publishes_application_work( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, failure: str, +) -> None: + fixture = server(tmp_path) + setattr(fixture, failure, True) + worker = await fixture.worker(monkeypatch) + assert await worker._run_workflow_task(fixture.task()) is None + assert all(not marker.exists() for marker in fixture.markers) + assert not any(name == "outcome" for name, _ in fixture.trace) + fixture.client.complete_workflow_task.assert_not_awaited() + fixture.client.fail_workflow_task.assert_not_awaited() + + +@pytest.mark.parametrize("failure", ["lost_outcome", "omit_second_outcome"]) +async def test_unknown_or_noncanonical_group_result_joins_siblings_and_never_reinvokes( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, failure: str, +) -> None: + fixture = server(tmp_path) + setattr(fixture, failure, True) + worker = await fixture.worker(monkeypatch) + assert await worker._run_workflow_task(fixture.task()) is None + assert sum(name == "prepare" for name, _ in fixture.trace) == 2 + for marker in fixture.markers: + await wait_for_exit(int(marker.read_text())) + assert not worker._prepared_local_activity_processes + fixture.client.complete_workflow_task.assert_not_awaited() + fixture.client.fail_workflow_task.assert_not_awaited() + + +async def test_cancellation_joins_both_blocked_callbacks_without_application_heartbeats_before_replay( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch, handler=blocked_without_python_progress) + task = fixture.task() + initial = capture(history=fixture.history, inputs=[str(tmp_path / "callback")]) + fixture.history = await worker._execute_prepared_local_activity_group(task, [], initial) + outcome = capture(history=fixture.history, inputs=[str(tmp_path / "callback")]) + # The blocked fixture accepts one marker, with no app heartbeat. + for call in outcome.prepared_local_activity_group.calls: + call.command.arguments = call.command.arguments[:1] + pending = asyncio.create_task(worker._execute_prepared_local_activity_group(task, fixture.history, outcome)) + try: + for marker in fixture.markers: + await wait_for_file(marker) + fixture.stop = True + with pytest.raises(PreparedCancellationObserved): + await asyncio.wait_for(pending, timeout=8) + for marker in fixture.markers: + await wait_for_exit(int(marker.read_text())) + assert sum(name == "acknowledge-cancellation" for name, _ in fixture.trace) == 2 + assert not any(name in {"heartbeat", "outcome"} for name, _ in fixture.trace) + assert not worker._prepared_local_activity_processes + finally: + pending.cancel() + await asyncio.gather(pending, return_exceptions=True) + + +async def test_cancel_during_second_admission_acks_unstarted_first_attempt_without_spawning( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch) + task = fixture.task() + fixture.history = await worker._execute_prepared_local_activity_group(task, [], capture()) + fixture.admission_cancelled = True + fixture.client.heartbeat_workflow_task.return_value = { + "task_id": "task", "lease_owner": "owner", "workflow_task_attempt": 3, "renewed": True, + "cancellation_request": {"request_id": "request-1", "requested_at": timestamp(), + "cleanup_deadline_at": timestamp(30), "history_refresh_page_token": "canonical-page"}, + } + with pytest.raises(LocalActivityExecutionAborted): + await worker._execute_prepared_local_activity_group(task, fixture.history, capture(history=fixture.history)) + assert all(not marker.exists() for marker in fixture.markers) + assert sum(name == "acknowledge-cancellation" for name, _ in fixture.trace) == 1 + + +async def test_group_capability_is_explicit_and_requires_the_installed_atomic_bridge( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + fixture = server(tmp_path) + worker = await fixture.worker(monkeypatch) + manifest = fixture.client.register_worker.await_args.kwargs["capability_manifest"] + assert manifest["prepared_local_activity_groups"]["implementation"] == "durable_atomic_all_admission" + capabilities = fixture.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"] + del capabilities["prepared_local_activity_groups"] + with pytest.raises(RuntimeError, match="installed atomic bridge"): + await worker._register() + + +def test_immutable_atomic_group_fixture_recovers_only_the_unfinished_original_sibling() -> None: + fixture = json.loads((Path(__file__).parent / "fixtures/replay_regressions/prepared-local-group-cold-results.json") + .read_text()) + outcome = capture(PreparedLocalGroupColdResultsWorkflow, fixture["history"][:-1], []) + group = outcome.prepared_local_activity_group + assert group is not None and group.committed and [call.sequence for call in group.calls] == [2] + assert group.calls[0].recover is True and group.commands == () + completed = capture(PreparedLocalGroupColdResultsWorkflow, fixture["history"], []) + assert completed.prepared_local_activity_group is None + assert completed.commands[0].result == ["fast-value", "fast-value"] diff --git a/tests/test_prepared_local_activity_worker.py b/tests/test_prepared_local_activity_worker.py index 1e49e0b..f25ec54 100644 --- a/tests/test_prepared_local_activity_worker.py +++ b/tests/test_prepared_local_activity_worker.py @@ -409,7 +409,7 @@ async def test_prepared_capability_requires_actual_bridge_and_never_advertises_g del info["worker_protocol"]["server_capabilities"]["prepared_local_activities"] with pytest.raises(RuntimeError, match="installed admission bridge"): await worker._register() - with pytest.raises(ValueError, match="no prepared local group consumer"): + with pytest.raises(ValueError, match="require prepared_local_activities"): Worker(server.client, task_queue="queue", capabilities=["prepared_local_activity_groups"]) diff --git a/tests/test_replay_regression_corpus.py b/tests/test_replay_regression_corpus.py index 64e1f22..a477770 100644 --- a/tests/test_replay_regression_corpus.py +++ b/tests/test_replay_regression_corpus.py @@ -220,6 +220,12 @@ def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] return [first, second] +@workflow.defn(name="tests.replay.prepared-local-group-cold-results") +class PreparedLocalGroupColdResultsWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + return (yield [ctx.local_activity("prepared.first", []), ctx.local_activity("prepared.second", [])]) + + @workflow.defn(name="tests.replay.cooperative-reopened-condition-cleanup") class CooperativeReopenedConditionCleanupWorkflow: def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] @@ -249,6 +255,7 @@ def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] GoldenVersionMarkerWorkflow, LocalActivityColdResultWorkflow, PreparedLocalColdResultsWorkflow, + PreparedLocalGroupColdResultsWorkflow, MessageStreamConsumerWorkflow, NestedParallelPathWorkflow, ParallelMetadataProducerWorkflow, From e256cbdb951cd9da48dde05ea272324740b74779 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 12:30:03 +0000 Subject: [PATCH 26/41] Retain prepared group refusals and mixed admission history --- src/durable_workflow/worker.py | 4 +++- .../test_prepared_local_activity.py | 18 +++++++++++++++--- 2 files changed, 18 insertions(+), 4 deletions(-) diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index b3d871c..56f58bf 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -1924,7 +1924,9 @@ async def _execute_prepared_local_activity_group( except PreparedCancellationObserved: await runner.acknowledge_stop() await self._renew_local_workflow_lease(task) - raise LocalActivityExecutionAborted("prepared group operation has an unknown or refused outcome") from error + raise LocalActivityExecutionAborted( + f"prepared group operation was refused: {error.reason() or 'unknown'}", + ) from error except (_WorkflowClaimDeferred, LocalActivityExecutionAborted): raise except Exception as error: diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index b54b4d3..975b249 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -14,7 +14,7 @@ import pytest from durable_workflow import Client, Worker, activity, workflow -from durable_workflow.errors import WorkflowCancelled +from durable_workflow.errors import ServerError, WorkflowCancelled from durable_workflow.workflow import replay from tests.integration.test_cooperative_cancellation import callback_gone, events, poll_claim, remote_marker @@ -127,7 +127,13 @@ def __init__(self, *args: Any, trace_path: Path, **kwargs: Any) -> None: self.trace_path = trace_path async def prepared_local_activity_operation(self, **kwargs: Any) -> dict[str, Any]: - receipt = await super().prepared_local_activity_operation(**kwargs) + try: + receipt = await super().prepared_local_activity_operation(**kwargs) + except ServerError as error: + with self.trace_path.open("a") as stream: + stream.write(json.dumps({"operation": kwargs["operation"], "sequence": kwargs["body"].get("sequence"), + "refused": error.reason()}) + "\n") + raise if kwargs["operation"] != "control" or receipt.get("active") is False: with self.trace_path.open("a") as stream: stream.write(json.dumps({"operation": kwargs["operation"], "receipt": receipt}) + "\n") @@ -173,7 +179,10 @@ async def test_prepared_nested_mixed_group_starts_peers_concurrently_and_cold_re ) -> None: queue = "py-prepared-group-" + uuid.uuid4().hex[:8] marker = str(tmp_path / "callback") - async with Client(server_url, token=server_token, namespace="default") as client: + trace_path = tmp_path / "group-receipts.jsonl" + async with TrackingPreparedClient( + server_url, token=server_token, namespace="default", trace_path=trace_path, + ) as client: worker = prepared_worker(client, queue) owner = asyncio.create_task(worker.run()) try: @@ -200,6 +209,9 @@ async def test_prepared_nested_mixed_group_starts_peers_concurrently_and_cold_re finally: await worker.stop() await asyncio.wait_for(owner, timeout=10) + if trace_path.exists(): + print("prepared nested mixed group receipts: " + trace_path.read_text()) + print("prepared nested mixed group final history: " + json.dumps(await events(handle))) @pytest.mark.parametrize("group", [False, True], ids=["sequential", "atomic-group"]) From f927f698aa2af0c882c0a701975fb485199153a1 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 12:47:45 +0000 Subject: [PATCH 27/41] Preserve prepared application heartbeat progress --- .../_prepared_local_activity.py | 2 +- .../test_prepared_local_activity.py | 7 ++++- tests/test_prepared_local_activity_worker.py | 28 +++++++++++++++++++ 3 files changed, 35 insertions(+), 2 deletions(-) diff --git a/src/durable_workflow/_prepared_local_activity.py b/src/durable_workflow/_prepared_local_activity.py index 856fa65..358128d 100644 --- a/src/durable_workflow/_prepared_local_activity.py +++ b/src/durable_workflow/_prepared_local_activity.py @@ -247,7 +247,7 @@ async def control(self, details: dict[str, Any] | None = None, *, heartbeat: boo _require(not self.shutdown.is_set(), "worker shutdown abandoned its prepared local claim") started = time.monotonic() receipt = await self.operation("heartbeat" if heartbeat else "control", - {"details": details} if heartbeat else {"renew_lease": True}) + {"progress": details or {}} if heartbeat else {"renew_lease": True}) self.attempt.validate_control(receipt, heartbeat=heartbeat) if receipt["active"]: self.attempt.accept_budget(receipt, started) diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index 975b249..24a6293 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -132,7 +132,8 @@ async def prepared_local_activity_operation(self, **kwargs: Any) -> dict[str, An except ServerError as error: with self.trace_path.open("a") as stream: stream.write(json.dumps({"operation": kwargs["operation"], "sequence": kwargs["body"].get("sequence"), - "refused": error.reason()}) + "\n") + "refused": error.reason(), "status": error.status, + "response": error.body}) + "\n") raise if kwargs["operation"] != "control" or receipt.get("active") is False: with self.trace_path.open("a") as stream: @@ -162,6 +163,8 @@ async def test_prepared_prefix_two_callbacks_and_cold_replay_use_canonical_histo kinds = [event["event_type"] for event in history] assert kinds.count("ActivityStarted") == kinds.count("ActivityCompleted") == 2 assert kinds.count("ActivityHeartbeatRecorded") == 2 + assert sorted(event["payload"]["progress"]["phase"] for event in history + if event["event_type"] == "ActivityHeartbeatRecorded") == ["first", "second"] assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 outcome = replay(PreparedSequentialWorkflow, history, [marker], run_id=handle.run_id or "", prepare_local_activities=True) @@ -197,6 +200,8 @@ async def test_prepared_nested_mixed_group_starts_peers_concurrently_and_cold_re assert kinds.count("ActivityScheduled") == kinds.count("ActivityStarted") == 2 assert kinds.count("ActivityCompleted") == 2 assert kinds.count("ActivityHeartbeatRecorded") == 2 + assert sorted(event["payload"]["progress"]["phase"] for event in history + if event["event_type"] == "ActivityHeartbeatRecorded") == ["first", "second"] assert kinds.count("TimerScheduled") == kinds.count("TimerFired") == 1 assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 completed = replay(PreparedGroupWorkflow, history, [marker], run_id=handle.run_id or "", diff --git a/tests/test_prepared_local_activity_worker.py b/tests/test_prepared_local_activity_worker.py index f25ec54..80feff8 100644 --- a/tests/test_prepared_local_activity_worker.py +++ b/tests/test_prepared_local_activity_worker.py @@ -59,6 +59,13 @@ def typed_callback(marker: str) -> dict[str, Any]: return {"value": b"\x00\xff", "attempt": activity.context().info.activity_attempt_id} +@activity.defn(name="prepared.progress") +async def progress_callback(marker: str) -> dict[str, Any]: + Path(marker).write_text(str(os.getpid())) + await activity.context().heartbeat({"phase": "processing", "count": 2}) + return {"value": b"\x00\xff", "attempt": activity.context().info.activity_attempt_id} + + def timestamp(seconds: float = 0) -> str: return (datetime.now(timezone.utc) + timedelta(seconds=seconds)).isoformat(timespec="microseconds").replace( "+00:00", "Z", @@ -301,6 +308,12 @@ async def operation(self, **kwargs: Any) -> dict[str, Any]: with pytest.raises(ProcessLookupError): os.kill(pid, 0) return {"acknowledged": True, "duplicate": False, "reason": None, "history_event_id": "joined"} + if name == "heartbeat": + assert body == {"progress": {"phase": "processing", "count": 2}} + payload = {**self.history[-1]["payload"], "progress": body["progress"]} + self.history.append({"id": "progress", "event_type": "ActivityHeartbeatRecorded", "payload": payload}) + return control(self.receipt, renewed=False, heartbeat_recorded=True, + heartbeat_history_event_id="progress") assert name == "outcome" if self.fail_outcome: raise TimeoutError("outcome acknowledgment lost") @@ -353,6 +366,21 @@ async def test_bad_admission_never_spawns_and_never_completes_or_fails_a_claim( server.client.fail_workflow_task.assert_not_awaited() +async def test_supervised_application_heartbeat_preserves_progress_in_canonical_history( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch, handler=progress_callback) + marker = tmp_path / "callback" + commands = await worker._run_workflow_task(server.task(marker)) + assert commands is not None and commands[0]["type"] == "complete_workflow" + progress = [event for event in server.history if event["event_type"] == "ActivityHeartbeatRecorded"] + assert len(progress) == 1 and progress[0]["payload"]["progress"] == {"phase": "processing", "count": 2} + assert server.trace.count("heartbeat") == 1 and server.trace.count("outcome") == 1 + await wait_for_exit(int(marker.read_text())) + server.client.fail_workflow_task.assert_not_awaited() + + @pytest.mark.parametrize("failure", ["fail_outcome", "omit_outcome_history"]) async def test_lost_or_noncanonical_outcome_abandons_without_reexecuting_callback( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, failure: str, From a6b243e84532e56d04d799e7b1e2c387874c5e3e Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 13:03:57 +0000 Subject: [PATCH 28/41] Use the canonical prepared heartbeat details envelope --- src/durable_workflow/_prepared_local_activity.py | 3 ++- tests/integration/test_prepared_local_activity.py | 4 ++-- tests/test_prepared_local_activity_worker.py | 4 ++-- 3 files changed, 6 insertions(+), 5 deletions(-) diff --git a/src/durable_workflow/_prepared_local_activity.py b/src/durable_workflow/_prepared_local_activity.py index 358128d..647f051 100644 --- a/src/durable_workflow/_prepared_local_activity.py +++ b/src/durable_workflow/_prepared_local_activity.py @@ -247,7 +247,8 @@ async def control(self, details: dict[str, Any] | None = None, *, heartbeat: boo _require(not self.shutdown.is_set(), "worker shutdown abandoned its prepared local claim") started = time.monotonic() receipt = await self.operation("heartbeat" if heartbeat else "control", - {"progress": details or {}} if heartbeat else {"renew_lease": True}) + {"progress": {"details": details} if details else {}} + if heartbeat else {"renew_lease": True}) self.attempt.validate_control(receipt, heartbeat=heartbeat) if receipt["active"]: self.attempt.accept_budget(receipt, started) diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index 24a6293..976c22c 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -163,7 +163,7 @@ async def test_prepared_prefix_two_callbacks_and_cold_replay_use_canonical_histo kinds = [event["event_type"] for event in history] assert kinds.count("ActivityStarted") == kinds.count("ActivityCompleted") == 2 assert kinds.count("ActivityHeartbeatRecorded") == 2 - assert sorted(event["payload"]["progress"]["phase"] for event in history + assert sorted(event["payload"]["progress"]["details"]["phase"] for event in history if event["event_type"] == "ActivityHeartbeatRecorded") == ["first", "second"] assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 outcome = replay(PreparedSequentialWorkflow, history, [marker], run_id=handle.run_id or "", @@ -200,7 +200,7 @@ async def test_prepared_nested_mixed_group_starts_peers_concurrently_and_cold_re assert kinds.count("ActivityScheduled") == kinds.count("ActivityStarted") == 2 assert kinds.count("ActivityCompleted") == 2 assert kinds.count("ActivityHeartbeatRecorded") == 2 - assert sorted(event["payload"]["progress"]["phase"] for event in history + assert sorted(event["payload"]["progress"]["details"]["phase"] for event in history if event["event_type"] == "ActivityHeartbeatRecorded") == ["first", "second"] assert kinds.count("TimerScheduled") == kinds.count("TimerFired") == 1 assert kinds.count("MemoUpserted") == kinds.count("WorkflowCompleted") == 1 diff --git a/tests/test_prepared_local_activity_worker.py b/tests/test_prepared_local_activity_worker.py index 80feff8..096032f 100644 --- a/tests/test_prepared_local_activity_worker.py +++ b/tests/test_prepared_local_activity_worker.py @@ -309,7 +309,7 @@ async def operation(self, **kwargs: Any) -> dict[str, Any]: os.kill(pid, 0) return {"acknowledged": True, "duplicate": False, "reason": None, "history_event_id": "joined"} if name == "heartbeat": - assert body == {"progress": {"phase": "processing", "count": 2}} + assert body == {"progress": {"details": {"phase": "processing", "count": 2}}} payload = {**self.history[-1]["payload"], "progress": body["progress"]} self.history.append({"id": "progress", "event_type": "ActivityHeartbeatRecorded", "payload": payload}) return control(self.receipt, renewed=False, heartbeat_recorded=True, @@ -375,7 +375,7 @@ async def test_supervised_application_heartbeat_preserves_progress_in_canonical_ commands = await worker._run_workflow_task(server.task(marker)) assert commands is not None and commands[0]["type"] == "complete_workflow" progress = [event for event in server.history if event["event_type"] == "ActivityHeartbeatRecorded"] - assert len(progress) == 1 and progress[0]["payload"]["progress"] == {"phase": "processing", "count": 2} + assert len(progress) == 1 and progress[0]["payload"]["progress"]["details"] == {"phase": "processing", "count": 2} assert server.trace.count("heartbeat") == 1 and server.trace.count("outcome") == 1 await wait_for_exit(int(marker.read_text())) server.client.fail_workflow_task.assert_not_awaited() From f7403886ce9abab2b32ee83656237d47efab6da6 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 15:16:22 +0000 Subject: [PATCH 29/41] Accept activity cancellation waits with explicit claim release --- src/durable_workflow/client.py | 8 +++-- src/durable_workflow/worker.py | 2 +- tests/test_cooperative_cancellation_client.py | 35 ++++++++++++------- 3 files changed, 29 insertions(+), 16 deletions(-) diff --git a/src/durable_workflow/client.py b/src/durable_workflow/client.py index 1afd6ff..5c85680 100644 --- a/src/durable_workflow/client.py +++ b/src/durable_workflow/client.py @@ -5228,8 +5228,12 @@ async def deliver_workflow_cancellation( if ( isinstance(result, dict) and result.get("delivered") is False - and delivery.call_kind in {"child", "parallel", "selection_handle"} - and result.get("reason") == "cancellation_waiting_for_child" + and ( + (result.get("reason") == "cancellation_waiting_for_child" + and delivery.call_kind in {"child", "parallel", "selection_handle"}) + or (result.get("reason") == "cancellation_waiting_for_activity" + and delivery.call_kind in {"activity", "local_activity", "parallel", "selection_handle"}) + ) and result.get("claim_released") is True and result.get("task_id") == task_id and all(result.get(field) is None for field in ( diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 56f58bf..99e1fab 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -186,7 +186,7 @@ class _CooperativeCancellationObserved(LocalActivityExecutionAborted): class _WorkflowClaimDeferred(LocalActivityExecutionAborted): - """Server parked the parent and released its claim until child cleanup ends.""" + """Server released this claim until cancellation acknowledgments resolve.""" class _InvalidLocalActivityReport(NonRetryableError): diff --git a/tests/test_cooperative_cancellation_client.py b/tests/test_cooperative_cancellation_client.py index 8925799..13f6dec 100644 --- a/tests/test_cooperative_cancellation_client.py +++ b/tests/test_cooperative_cancellation_client.py @@ -200,13 +200,17 @@ async def test_delivery_sends_owner_attempt_and_authored_boundary( @pytest.mark.asyncio -@pytest.mark.parametrize("kind", ["child", "parallel", "selection_handle"]) -async def test_pending_child_requires_explicit_claim_release( - client: Client, monkeypatch: pytest.MonkeyPatch, kind: str, +@pytest.mark.parametrize("kind,reason", [ + (kind, "cancellation_waiting_for_child") for kind in ("child", "parallel", "selection_handle") +] + [ + (kind, "cancellation_waiting_for_activity") for kind in ("activity", "local_activity", "parallel", "selection_handle") +]) +async def test_pending_cancellation_requires_explicit_claim_release( + client: Client, monkeypatch: pytest.MonkeyPatch, kind: str, reason: str, ) -> None: monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") options = {"operation_sequence": 1} if kind == "selection_handle" else {} - pending = pending_delivery_response() + pending = pending_delivery_response(reason=reason) with patch.object(client._http, "request", new_callable=AsyncMock, return_value=response(pending)): result = await client.deliver_workflow_cancellation( task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, @@ -216,41 +220,46 @@ async def test_pending_child_requires_explicit_claim_release( @pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["child", "activity"]) @pytest.mark.parametrize("change", [ {"claim_released": False}, {"claim_released": "true"}, {"claim_released": None}, - {"task_id": "other"}, {"reason": "other"}, {"request_id": "original-request"}, + {"task_id": "other"}, {"reason": "other"}, {"reason": []}, {"request_id": "original-request"}, {"sequence": 3}, {"call_kind": "child"}, {"sequence_span": 1}, {"operation_sequence": 1}, {"operation_sequence_span": 1}, {"delivered": 0}, ]) -async def test_malformed_pending_child_ack_is_rejected( - client: Client, monkeypatch: pytest.MonkeyPatch, change: dict[str, Any], +async def test_malformed_pending_cancellation_ack_is_rejected( + client: Client, monkeypatch: pytest.MonkeyPatch, change: dict[str, Any], kind: str, ) -> None: monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") with ( patch.object(client._http, "request", new_callable=AsyncMock, - return_value=response(pending_delivery_response(**change))), + return_value=response(pending_delivery_response(reason="cancellation_waiting_for_" + kind) | change)), pytest.raises(ServerError) as error, ): await client.deliver_workflow_cancellation( task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, - request_id="original-request", sequence=3, call_kind="child", + request_id="original-request", sequence=3, call_kind=kind, ) assert error.value.reason() == "invalid_cooperative_cancellation_delivery" @pytest.mark.asyncio -async def test_pending_child_reply_cannot_release_an_unrelated_timer_claim( - client: Client, monkeypatch: pytest.MonkeyPatch, +@pytest.mark.parametrize("kind,reason", [ + ("timer", "cancellation_waiting_for_child"), ("timer", "cancellation_waiting_for_activity"), + ("child", "cancellation_waiting_for_activity"), ("activity", "cancellation_waiting_for_child"), +]) +async def test_pending_reply_cannot_release_an_unrelated_claim( + client: Client, monkeypatch: pytest.MonkeyPatch, kind: str, reason: str, ) -> None: monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") with ( patch.object(client._http, "request", new_callable=AsyncMock, - return_value=response(pending_delivery_response())), + return_value=response(pending_delivery_response(reason=reason))), pytest.raises(ServerError), ): await client.deliver_workflow_cancellation( task_id="task/1", lease_owner="worker-1", workflow_task_attempt=2, - request_id="original-request", sequence=3, call_kind="timer", + request_id="original-request", sequence=3, call_kind=kind, ) From bf931959af5cc6615cb42ccd0ed279e7b3d33024 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 15:20:21 +0000 Subject: [PATCH 30/41] Format pending cancellation response cases --- tests/test_cooperative_cancellation_client.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/test_cooperative_cancellation_client.py b/tests/test_cooperative_cancellation_client.py index 13f6dec..11a3f81 100644 --- a/tests/test_cooperative_cancellation_client.py +++ b/tests/test_cooperative_cancellation_client.py @@ -203,7 +203,8 @@ async def test_delivery_sends_owner_attempt_and_authored_boundary( @pytest.mark.parametrize("kind,reason", [ (kind, "cancellation_waiting_for_child") for kind in ("child", "parallel", "selection_handle") ] + [ - (kind, "cancellation_waiting_for_activity") for kind in ("activity", "local_activity", "parallel", "selection_handle") + (kind, "cancellation_waiting_for_activity") + for kind in ("activity", "local_activity", "parallel", "selection_handle") ]) async def test_pending_cancellation_requires_explicit_claim_release( client: Client, monkeypatch: pytest.MonkeyPatch, kind: str, reason: str, @@ -231,9 +232,10 @@ async def test_malformed_pending_cancellation_ack_is_rejected( client: Client, monkeypatch: pytest.MonkeyPatch, change: dict[str, Any], kind: str, ) -> None: monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") + pending = pending_delivery_response(reason="cancellation_waiting_for_" + kind) | change with ( patch.object(client._http, "request", new_callable=AsyncMock, - return_value=response(pending_delivery_response(reason="cancellation_waiting_for_" + kind) | change)), + return_value=response(pending)), pytest.raises(ServerError) as error, ): await client.deliver_workflow_cancellation( From de0c476760524dcac3df35572bb7bd19d818a2d0 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 16:44:25 +0000 Subject: [PATCH 31/41] Preserve remote activity cancellation policies through Python replay --- docs/cooperative-cancellation-design.md | 38 +++++ src/durable_workflow/cancellation.py | 15 +- src/durable_workflow/worker.py | 18 +++ src/durable_workflow/workflow.py | 62 +++++++- .../activity-cancellation-policy-changed.json | 13 ++ .../test_cooperative_cancellation.py | 94 +++++++++++- tests/test_activity_cancellation_policies.py | 136 ++++++++++++++++++ 7 files changed, 367 insertions(+), 9 deletions(-) create mode 100644 tests/fixtures/replay_regressions/activity-cancellation-policy-changed.json create mode 100644 tests/test_activity_cancellation_policies.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index a529ecb..7b2e150 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -70,6 +70,44 @@ conflicting history fail replay. A worker without the negotiated cooperation capability refuses these choices before completion, reporting its identity and the required protocol. Server also checks the immutable task claim and backend. +## Remote Activity policies + +`ctx.schedule_activity()` accepts `cancellation_policy` as a `CancellationPolicy` +enum or its portable string value. Both command encoders preserve it. Omission +keeps the historical `TRY_CANCEL` behavior and wire shape. + +```python +yield ctx.schedule_activity( + "rust.remote-work", [], + cancellation_policy=CancellationPolicy.WAIT_CANCELLATION_COMPLETED, +) +``` + +`TRY_CANCEL` requests cancellation and continues without waiting for the stop +receipt. `WAIT_CANCELLATION_COMPLETED` delays workflow delivery until the +original remote attempt's physical stop acknowledgment is recorded. `ABANDON` +leaves the Activity independent after parent cancellation. It requires a finite +positive integer `schedule_to_close_timeout`, whose original deadline bounds +the independent work. It does not extend the parent's cleanup budget. + +Every explicit remote policy requires negotiated cooperation, protocol 1.20 +and a compatible installed backend. Missing worker capability is diagnosed +before submission. Server also checks the original immutable claim. Local +Activity policy authoring remains unavailable while its lifetime and admission +contract are unfinished. + +Replay compares authored policy with original canonical history for ordinary, +parallel and selection calls and cancellation delivery. Later events that omit +the policy retain the original value. Unknown, conflicting and changed policies +fail replay explicitly. Historical histories without a policy retain Try. + +The connected Source tests include explicit Try and Wait with both async and +sync callbacks and no application heartbeats. Wait checks physical stop receipt +ordering before workflow delivery. Bounded Abandon checks that a callback +survives parent closure, completes under its original total lifetime and cannot +publish a second outcome or reopen the cancelled parent. The exact published +mixed-language gate remains required. + ## Remote callback-stop transport `Client.acknowledge_activity_cancellation()` reports the original task, activity diff --git a/src/durable_workflow/cancellation.py b/src/durable_workflow/cancellation.py index aeb0beb..23d3e7a 100644 --- a/src/durable_workflow/cancellation.py +++ b/src/durable_workflow/cancellation.py @@ -12,7 +12,7 @@ class CancellationPolicy(str, Enum): - """Cancellation at an awaiting child call. Abandon preserves the legacy default.""" + """Cancellation at an awaiting operation. Activities default to Try, children to Abandon.""" TRY_CANCEL = "try_cancel" WAIT_CANCELLATION_COMPLETED = "wait_cancellation_completed" @@ -48,6 +48,19 @@ def _canonical_child_policies(options: Mapping[str, Any]) -> dict[str, str]: return policies +def _canonical_activity_policy(value: Any) -> str | None: + if value is None: + return None + if isinstance(value, Enum) and not isinstance(value, CancellationPolicy): + raise ValueError("remote activity cancellation_policy must be a supported policy") + if not isinstance(value, str): + raise ValueError("remote activity cancellation_policy must be a supported policy") + try: + return CancellationPolicy(value).value + except ValueError as error: + raise ValueError("remote activity cancellation_policy must be a supported policy") from error + + def _text(snapshot: Mapping[str, Any], key: str) -> str: value = snapshot.get(key) if not isinstance(value, str) or not value.strip(): diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 99e1fab..19b4d5b 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -98,6 +98,7 @@ RecordLocalActivity, RecordSideEffect, ReplayOutcome, + ScheduleActivity, StartChildWorkflow, UpsertMemo, apply_update, @@ -2470,6 +2471,23 @@ def execute_local(command: RecordLocalActivity) -> Any: log.warning("failed to report workflow memo capability failure: %s", failure_error) return None + if not self._cooperative_cancellation_supported and any( + isinstance(command, ScheduleActivity) and command.cancellation_policy is not None + for command in workflow_commands + ): + message = ( + f"activity_cancellation_policy_not_supported: Python worker {self.worker_id} requires " + "cooperative_cancellation capability, worker protocol 1.20 and a compatible Server/Native backend" + ) + try: + await self.client.fail_workflow_task( + task_id=task_id, lease_owner=self.worker_id, workflow_task_attempt=attempt, + message=message, failure_type="RuntimeCapabilityUnsupported", + ) + except Exception as failure_error: + log.warning("failed to report activity cancellation policy capability failure: %s", failure_error) + return None + if not self._cooperative_cancellation_supported and any( isinstance(command, StartChildWorkflow) and ( command.parent_close_policy == "request_cancellation" diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 78d0a02..0bc487b 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -34,7 +34,13 @@ from . import serializer from ._cooperative_cancellation import CancellationDelivery, read_cancellation_history -from .cancellation import CancellationContext, CancellationPolicy, ParentClosePolicy, _canonical_child_policies +from .cancellation import ( + CancellationContext, + CancellationPolicy, + ParentClosePolicy, + _canonical_activity_policy, + _canonical_child_policies, +) from .client import WorkflowStreamAppendItem from .errors import ( ActivityFailed, @@ -471,6 +477,7 @@ class ScheduleActivity: schedule_to_close_timeout: int | None = None heartbeat_timeout: int | None = None worker_session: WorkerSessionOptions | None = None + cancellation_policy: str | CancellationPolicy | None = None _parallel_group_path: list[dict[str, Any]] | None = field( default=None, init=False, @@ -478,6 +485,18 @@ class ScheduleActivity: compare=False, ) + def __post_init__(self) -> None: + self._validate_cancellation_policy() + + def _validate_cancellation_policy(self) -> None: + self.cancellation_policy = _canonical_activity_policy(self.cancellation_policy) + if self.cancellation_policy == CancellationPolicy.ABANDON.value and ( + not isinstance(self.schedule_to_close_timeout, int) + or isinstance(self.schedule_to_close_timeout, bool) + or self.schedule_to_close_timeout < 1 + ): + raise ValueError("remote activity Abandon requires a finite positive schedule_to_close_timeout") + def to_server_command( self, task_queue: str, @@ -488,6 +507,7 @@ def to_server_command( external_storage: ExternalStorageDriver | None = None, external_storage_threshold_bytes: int | None = None, ) -> dict[str, Any]: + self._validate_cancellation_policy() self._validate_timeouts() command: dict[str, Any] = { @@ -524,6 +544,8 @@ def to_server_command( command["heartbeat_timeout"] = self.heartbeat_timeout if self.worker_session is not None: command["worker_session"] = self.worker_session.to_wire() + if self.cancellation_policy is not None: + command["cancellation_policy"] = _canonical_activity_policy(self.cancellation_policy) return command def _validate_timeouts(self) -> None: @@ -1438,6 +1460,7 @@ def commands_to_server_commands( for command in commands: if isinstance(command, ScheduleActivity): + command._validate_cancellation_policy() queue = command.queue or task_queue server_command: dict[str, Any] = { "type": "schedule_activity", @@ -1471,6 +1494,8 @@ def commands_to_server_commands( server_command["heartbeat_timeout"] = command.heartbeat_timeout if command.worker_session is not None: server_command["worker_session"] = command.worker_session.to_wire() + if command.cancellation_policy is not None: + server_command["cancellation_policy"] = _canonical_activity_policy(command.cancellation_policy) server_commands.append(server_command) continue @@ -1930,6 +1955,7 @@ def schedule_activity( schedule_to_close_timeout: int | None = None, heartbeat_timeout: int | None = None, worker_session: WorkerSessionOptions | None = None, + cancellation_policy: str | CancellationPolicy | None = None, ) -> ScheduleActivity: return ScheduleActivity( activity_type=activity_type, @@ -1941,6 +1967,7 @@ def schedule_activity( schedule_to_close_timeout=schedule_to_close_timeout, heartbeat_timeout=heartbeat_timeout, worker_session=worker_session, + cancellation_policy=cancellation_policy, ) def local_activity( @@ -3545,6 +3572,13 @@ def _recorded_detail_mismatch(command: Any, step: _RecordedStep) -> str | None: f"Recorded activity_type {recorded!r}, but current workflow " f"scheduled {command.activity_type!r}." ) + actual_policy = _canonical_activity_policy(command.cancellation_policy) or "try_cancel" + recorded_policy = step.details.get("cancellation_policy", "try_cancel") + if recorded_policy != actual_policy: + return ( + f"activity_cancellation_policy_changed: recorded {recorded_policy!r}, " + f"but current workflow requested {actual_policy!r}." + ) elif isinstance(command, RecordLocalActivity): if step.details.get("execution_mode") != "local": return "Recorded remote activity cannot replay as a local activity." @@ -3756,6 +3790,7 @@ def _replay_state( event_types_by_sequence: dict[int, list[str]] = {} details_by_sequence: dict[int, dict[str, Any]] = {} child_policies_by_sequence: dict[int, dict[str, str]] = {} + activity_policies_by_sequence: dict[int, str] = {} resolved_sequences: set[int] = set() condition_wait_ids_by_sequence: dict[int, str] = {} selected_condition_wait_ids_by_sequence: dict[int, str] = {} @@ -3777,6 +3812,29 @@ def _replay_state( # envelopes into the public non-determinism diagnostic. recorded_details = {} details_by_sequence.setdefault(sequence, {}).update(recorded_details) + if event_type in ( + "ActivityScheduled", "ActivityStarted", "ActivityCompleted", + "ActivityFailed", "ActivityTimedOut", "ActivityCancelled", + ): + policy = activity_policies_by_sequence.get(sequence) + snapshot = payload.get("activity") + for source in (payload, snapshot if isinstance(snapshot, Mapping) else {}): + if "cancellation_policy" not in source: + continue + incoming = source["cancellation_policy"] + if not isinstance(incoming, str) or incoming not in tuple(item.value for item in CancellationPolicy): + raise NonDeterministicReplayError( + sequence, "activity", [event_type], + detail="invalid_activity_cancellation_policy_history: expected a supported policy.", + ) + if policy is not None and policy != incoming: + raise NonDeterministicReplayError( + sequence, "activity", [event_type], + detail="activity_cancellation_policy_history_conflict: policy changed between events.", + ) + policy = incoming + activity_policies_by_sequence[sequence] = policy or "try_cancel" + details_by_sequence[sequence]["cancellation_policy"] = activity_policies_by_sequence[sequence] if event_type in ( "ChildWorkflowScheduled", "ChildRunStarted", "ChildRunCompleted", "ChildRunFailed", "ChildRunCancelled", "ChildRunTerminated", @@ -3994,6 +4052,8 @@ def _recorded_step( details.update(_recorded_step_details(payload)) if shape == "child workflow": details.update(child_policies_by_sequence.get(workflow_sequence, {})) + if shape == "activity": + details["cancellation_policy"] = activity_policies_by_sequence.get(workflow_sequence, "try_cancel") return _RecordedStep( workflow_sequence=workflow_sequence, shape=shape, diff --git a/tests/fixtures/replay_regressions/activity-cancellation-policy-changed.json b/tests/fixtures/replay_regressions/activity-cancellation-policy-changed.json new file mode 100644 index 0000000..4f781ac --- /dev/null +++ b/tests/fixtures/replay_regressions/activity-cancellation-policy-changed.json @@ -0,0 +1,13 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "python-activity-cancellation-policy-changed", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": {"type": "golden.single-activity", "input": ["Ada"], "payload_codec": "avro"}, + "history": [ + {"event_type": "ActivityScheduled", "payload": {"sequence": 1, "activity": {"type": "golden.greet", "cancellation_policy": "wait_cancellation_completed"}}} + ], + "expected": {"command_sequence": []}, + "expected_replay_error": {"type": "NonDeterministicReplayError", "message_contains": "activity_cancellation_policy_changed", "workflow_sequence": 1} +} diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 0bf12b1..3a76f26 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -16,7 +16,7 @@ import pytest -from durable_workflow import Client, Worker, activity, workflow +from durable_workflow import CancellationPolicy, Client, Worker, activity, serializer, workflow from durable_workflow.client import WorkflowHandle from durable_workflow.errors import ServerError, WorkflowCancelled from durable_workflow.worker import _poll_capacity_delay @@ -38,12 +38,15 @@ async def cooperative_runtime(server_url: str, server_token: str, monkeypatch: p @workflow.defn(name="tests.python-cooperative-cleanup") class CooperativeCleanupWorkflow: - def run(self, ctx: Any, kind: str) -> Any: + def run(self, ctx: Any, kind: str, cancellation_policy: str | None = None) -> Any: try: if kind == "local": yield ctx.local_activity("tests.python-cooperative-work", []) elif kind == "remote": - yield ctx.schedule_activity("tests.python-cooperative-work", []) + yield ctx.schedule_activity( + "tests.python-cooperative-work", [], cancellation_policy=cancellation_policy, + schedule_to_close_timeout=60 if cancellation_policy is not None else None, + ) else: yield ctx.start_timer(300) except WorkflowCancelled as error: @@ -141,12 +144,17 @@ async def heartbeat_worker(self, **kwargs: Any) -> Any: class AsyncRemoteQualification: marker: str user_heartbeat: bool = False + duration_seconds: int | None = None - async def __call__(self) -> None: + async def __call__(self) -> str | None: if self.user_heartbeat: await activity.context().heartbeat({"qualification": "remote-in-flight"}) write_remote_marker(self.marker) + if self.duration_seconds is not None: + await asyncio.sleep(self.duration_seconds) + return "independent-completion" await asyncio.Event().wait() + return None @dataclass(frozen=True) @@ -205,9 +213,13 @@ async def receipt() -> dict[str, Any]: @pytest.mark.parametrize("handler_kind", ["async", "sync"]) -@pytest.mark.parametrize("user_heartbeat", [False, True]) +@pytest.mark.parametrize("user_heartbeat,policy", [ + (False, None), (True, None), (False, CancellationPolicy.TRY_CANCEL), + (False, CancellationPolicy.WAIT_CANCELLATION_COMPLETED), +]) async def test_actual_remote_worker_stops_callbacks_and_reports_original_cancellation( - server_url: str, server_token: str, handler_kind: str, user_heartbeat: bool, tmp_path: Path, + server_url: str, server_token: str, handler_kind: str, user_heartbeat: bool, policy: CancellationPolicy | None, + tmp_path: Path, ) -> None: queue = f"py-cooperative-owner-{uuid.uuid4().hex[:8]}" marker = tmp_path / "remote" @@ -220,7 +232,8 @@ async def test_actual_remote_worker_stops_callbacks_and_reports_original_cancell running = asyncio.create_task(worker.run()) try: handle = await client.start_workflow( - workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["remote"], + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, + input=["remote", policy.value if policy is not None else None], ) entered = await remote_marker(marker) await asyncio.wait_for(client.owner_heartbeat.wait(), timeout=15) @@ -253,6 +266,12 @@ async def test_actual_remote_worker_stops_callbacks_and_reports_original_cancell assert payload[key] == proof[key] delivery = [event for event in history if event["event_type"] == "CooperativeCancellationDelivered"][0] assert delivery["payload"]["call_kind"] == "activity" + if policy is not None: + scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"][0] + assert scheduled["payload"]["activity"]["cancellation_policy"] == policy.value + if policy == CancellationPolicy.WAIT_CANCELLATION_COMPLETED: + kinds = [event["event_type"] for event in history] + assert kinds.index("ActivityCancellationAcknowledged") < kinds.index("CooperativeCancellationDelivered") assert len([event for event in history if event["event_type"] == "ActivityCancelled"]) == 1 progress = [event for event in history if event["event_type"] == "ActivityHeartbeatRecorded" and event["payload"].get("activity_type") == "tests.python-cooperative-work"] @@ -270,6 +289,67 @@ async def test_actual_remote_worker_stops_callbacks_and_reports_original_cancell await asyncio.wait_for(running, timeout=10) +async def test_bounded_remote_abandon_completes_after_parent_cancellation( + server_url: str, server_token: str, tmp_path: Path, +) -> None: + queue = f"py-cooperative-abandon-{uuid.uuid4().hex[:8]}" + marker = tmp_path / "remote" + async with ObservedOwnerClient(server_url, token=server_token, namespace="default") as client: + worker = candidate_worker(client, queue, max_concurrent_activity_tasks=1) + worker.activities["tests.python-cooperative-work"] = AsyncRemoteQualification(str(marker), duration_seconds=25) + running = asyncio.create_task(worker.run()) + try: + handle = await client.start_workflow( + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, + input=["remote", CancellationPolicy.ABANDON.value], + ) + entered = await remote_marker(marker) + await asyncio.wait_for(client.owner_heartbeat.wait(), timeout=15) + fence = {key: entered[key] for key in ("task_id", "activity_attempt_id", "lease_owner")} + accepted = await handle.request_cancellation(cleanup_timeout_seconds=30) + duplicate = await handle.request_cancellation(cleanup_timeout_seconds=300) + assert duplicate["duplicate"] is True + assert duplicate["cancellation_request"] == accepted["cancellation_request"] + history = await assert_cancelled_cleanup(handle, accepted["cancellation_request"]["request_id"]) + os.kill(entered["callback_pid"], 0) + kinds = [event["event_type"] for event in history] + assert "ActivityCancelled" not in kinds + assert "ActivityCancellationAcknowledged" not in kinds + scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"][0] + assert scheduled["payload"]["activity"]["cancellation_policy"] == "abandon" + total_deadline = scheduled["payload"]["activity"]["schedule_to_close_deadline_at"] + assert isinstance(total_deadline, str) + status = await client.activity_task_status(**fence) + assert status["can_continue"] is True + until = time.monotonic() + 35 + while True: + history = await events(handle) + completed = [event for event in history if event["event_type"] == "ActivityCompleted" + and event["payload"].get("activity_type") == "tests.python-cooperative-work"] + if completed or time.monotonic() >= until: + break + await asyncio.sleep(0.1) + assert len(completed) == 1 + assert serializer.decode_envelope(completed[0]["payload"]["result"]) == "independent-completion" + assert completed[0]["payload"]["activity"]["schedule_to_close_deadline_at"] == total_deadline + assert completed[0]["payload"]["activity_attempt_id"] == fence["activity_attempt_id"] + kinds = [event["event_type"] for event in history] + assert kinds.count("WorkflowCancelled") == 1 + assert "WorkflowCompleted" not in kinds + assert "ActivityCancellationAcknowledged" not in kinds + with pytest.raises(ServerError) as completion: + await client.complete_activity_task(**fence, result="stale") + assert completion.value.status == 409 + with pytest.raises(ServerError) as failure: + await client.fail_activity_task(**fence, message="stale", failure_type="LateQualification") + assert failure.value.status == 409 + assert await events(handle) == history + print(f"Bounded remote Abandon: {json.dumps([entered, accepted, completed])}") + finally: + await worker.stop() + await asyncio.wait_for(running, timeout=10) + + @pytest.mark.parametrize("handler_kind", ["async", "sync"]) async def test_actual_remote_shutdown_joins_callback_before_replacement_cleanup( server_url: str, server_token: str, handler_kind: str, tmp_path: Path, diff --git a/tests/test_activity_cancellation_policies.py b/tests/test_activity_cancellation_policies.py new file mode 100644 index 0000000..f08dfdb --- /dev/null +++ b/tests/test_activity_cancellation_policies.py @@ -0,0 +1,136 @@ +from __future__ import annotations + +from typing import Any + +import pytest + +from durable_workflow import CancellationPolicy, ParentClosePolicy, serializer +from durable_workflow import workflow as workflow_module +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.worker import Worker +from durable_workflow.workflow import ( + CompleteWorkflow, + ScheduleActivity, + WorkflowContext, + commands_to_server_commands, + replay, +) +from tests.test_cooperative_cancellation import marker, request +from tests.test_cooperative_cancellation_worker import ClaimServer, claimed_task + + +def workflow(options: dict[str, Any], mode: str = "sequential") -> type: + class PolicyWorkflow: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + first = ctx.schedule_activity("work", [], **options) + try: + if mode == "parallel": + return (yield [first]) + if mode == "selection": + return (yield ctx.select({"work": first})) + return (yield first) + except WorkflowCancelled as cancelled: + return cancelled.request_id + + return PolicyWorkflow + + +def event(kind: str, policy: Any = None) -> dict[str, Any]: + activity = {"type": "work"} + if policy is not None: + activity["cancellation_policy"] = policy + return {"event_type": kind, "payload": {"sequence": 1, "activity": activity}} + + +def test_changed_activity_policy_is_rejected_during_replay() -> None: + with pytest.raises(NonDeterministicReplayError, match="activity_cancellation_policy_changed"): + replay(workflow({}), [event("ActivityScheduled", "wait_cancellation_completed")], []) + + +@pytest.mark.parametrize("policy", list(CancellationPolicy)) +def test_typed_and_string_policies_match_both_encoders_and_replay(policy: CancellationPolicy) -> None: + command = ScheduleActivity("work", [], schedule_to_close_timeout=60, cancellation_policy=policy) + direct = command.to_server_command("queue") + assert direct == commands_to_server_commands([command], "queue")[0] + assert direct == ScheduleActivity( + "work", [], schedule_to_close_timeout=60, cancellation_policy=policy.value, + ).to_server_command("queue") + assert type(direct["cancellation_policy"]) is str + history = [event("ActivityScheduled", policy.value), event("ActivityStarted"), event("ActivityCompleted")] + history[-1]["payload"]["result"] = serializer.envelope("recorded") + assert replay(workflow({"cancellation_policy": policy, "schedule_to_close_timeout": 60}), history, []).commands == [ + CompleteWorkflow("recorded"), + ] + + +def test_omitted_policy_preserves_wire_and_historical_try_default() -> None: + assert "cancellation_policy" not in ScheduleActivity("work", []).to_server_command("queue") + history = [event("ActivityCompleted")] + history[0]["payload"]["result"] = serializer.envelope("recorded") + assert replay(workflow({}), history, []).commands == [CompleteWorkflow("recorded")] + assert replay(workflow({"cancellation_policy": CancellationPolicy.TRY_CANCEL}), history, []).commands == [ + CompleteWorkflow("recorded"), + ] + + +@pytest.mark.parametrize("mode", ["sequential", "parallel", "selection"]) +@pytest.mark.parametrize("delivered", [False, True]) +def test_changed_policy_fails_group_matching_and_cancellation_delivery(mode: str, delivered: bool) -> None: + options = {"cancellation_policy": CancellationPolicy.WAIT_CANCELLATION_COMPLETED} + commands = commands_to_server_commands(replay(workflow(options, mode), [], []).commands, "queue") + history = [{"event_type": "ActivityScheduled", "payload": {"sequence": 1, **command}} for command in commands] + if delivered: + history += [request(), marker(1, "activity" if mode == "sequential" else "parallel", sequence_span=1)] + assert replay(workflow(options, mode), history, [], run_id="run-1").commands == [CompleteWorkflow("request-1")] + with pytest.raises(NonDeterministicReplayError, match="activity_cancellation_policy_changed"): + replay(workflow({"cancellation_policy": CancellationPolicy.TRY_CANCEL}, mode), history, [], run_id="run-1") + + +@pytest.mark.parametrize("value", ["unknown", "", True, 1, [], ParentClosePolicy.ABANDON]) +def test_invalid_policy_is_refused_before_suspension(value: Any) -> None: + with pytest.raises(ValueError, match="cancellation_policy"): + ScheduleActivity("work", [], cancellation_policy=value) + + +@pytest.mark.parametrize("timeout", [None, 0, -1, 1.5, "60", True]) +def test_abandon_requires_finite_total_timeout(timeout: Any) -> None: + with pytest.raises(ValueError, match="schedule_to_close_timeout"): + ScheduleActivity("work", [], cancellation_policy=CancellationPolicy.ABANDON, schedule_to_close_timeout=timeout) + + +@pytest.mark.parametrize("kind", ["ActivityScheduled", "ActivityCompleted"]) +@pytest.mark.parametrize("value", [None, "unknown", []]) +def test_malformed_canonical_history_is_rejected(kind: str, value: Any) -> None: + item = event(kind) + item["payload"]["activity"]["cancellation_policy"] = value + with pytest.raises(NonDeterministicReplayError, match="invalid_activity_cancellation_policy_history"): + replay(workflow({}), [item], []) + + +def test_conflicting_history_is_rejected() -> None: + with pytest.raises(NonDeterministicReplayError, match="activity_cancellation_policy_history_conflict"): + replay(workflow({}), [event("ActivityScheduled", "try_cancel"), event("ActivityStarted", "abandon")], []) + + +@pytest.mark.parametrize("protocol", ["1.19", "1.20"]) +@pytest.mark.parametrize("policy", list(CancellationPolicy)) +async def test_worker_without_opt_in_refuses_explicit_policy( + monkeypatch: pytest.MonkeyPatch, protocol: str, policy: CancellationPolicy, +) -> None: + monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", protocol) + cls = workflow_module.defn(name="activity-policy-worker")(workflow({ + "cancellation_policy": policy, "schedule_to_close_timeout": 60, + })) + server = ClaimServer(history=[]) + worker = Worker(server.client, task_queue="queue", worker_id="cooperative-worker", workflows=[cls]) + await worker._register() + result = await worker._run_workflow_task({ + **claimed_task(observed=False), "workflow_type": "activity-policy-worker", "arguments": serializer.envelope([]), + }) + assert result is None + server.client.complete_workflow_task.assert_not_awaited() + failure = server.client.fail_workflow_task.await_args.kwargs + assert failure["failure_type"] == "RuntimeCapabilityUnsupported" + assert "activity_cancellation_policy_not_supported" in failure["message"] + assert "cooperative-worker" in failure["message"] + assert "1.20" in failure["message"] From 8c1e937d2bab234a8a9cd1c1cea3aed7265aa9db Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 21:16:16 +0000 Subject: [PATCH 32/41] Add deterministic cancellation remaining time during replay --- docs/cooperative-cancellation-design.md | 33 ++- src/durable_workflow/cancellation.py | 30 +- src/durable_workflow/workflow.py | 96 ++++++- .../test_cooperative_cancellation.py | 29 +- tests/test_cancellation_remaining_time.py | 265 ++++++++++++++++++ 5 files changed, 437 insertions(+), 16 deletions(-) create mode 100644 tests/test_cancellation_remaining_time.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index 7b2e150..db5aee9 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -211,11 +211,34 @@ the same original root, delivery event and immutable deadline. ## Remaining qualification -Rust policy parity, portable activity policies, nested scopes and deterministic -remaining-time helpers still need completion. Remaining time must use the -replayed workflow clock. Do not subtract the host clock from the deadline in -workflow code. The runtime continues enforcing the original deadline and fencing -task and activity ownership. +### Python deadline and remaining time + +The Source `CancellationContext.deadline` is the original immutable cleanup +deadline. `remaining()` returns fractional seconds left at the replay boundary +consumed by workflow code, clamped to zero. Committed cancellation delivery sets +the initial clock. Blocking activity, prepared local activity, child, timer, +condition, selection and awaited-handle outcomes advance it. A selection uses +its committed winner marker. A group failure excludes later sibling outcomes. +Recorded clock skew cannot increase an already consumed budget. + +Synchronous side effects, version markers, memo updates and inline local +callbacks preserve that clock because they return before their results are +persisted on first execution. Their later history timestamps cannot change +the same authored decision during cold replay. `WorkflowContext.now()` retains +its existing start-time contract. + +Only the active replay that delivered the context can use its remaining-time +clock. Detached metadata and calls after replay ends fail explicitly. Missing +or invalid boundary timestamps also fail, without a host-time fallback. +The runtime supervisor independently enforces the actual deadline and task +ownership even when authoring code cannot run. + +The connected process-loss scenario records remaining time before SIGKILL, +requires the same value in the replacement worker, then checks the value after +prepared cleanup against its committed history timestamp and original deadline. + +Rust helpers, explicit local operation policies, nested scopes and competitive +qualification still need completion. Connected qualification must cover the PHP parent, Python child, Rust remote activity and PHP local activity together. Callbacks must stop without application diff --git a/src/durable_workflow/cancellation.py b/src/durable_workflow/cancellation.py index 23d3e7a..7edb275 100644 --- a/src/durable_workflow/cancellation.py +++ b/src/durable_workflow/cancellation.py @@ -3,8 +3,8 @@ from __future__ import annotations import re -from collections.abc import Mapping -from dataclasses import dataclass +from collections.abc import Callable, Mapping +from dataclasses import dataclass, field, replace from datetime import datetime, timezone from enum import Enum from types import MappingProxyType @@ -30,21 +30,21 @@ class ParentClosePolicy(str, Enum): def _canonical_child_policies(options: Mapping[str, Any]) -> dict[str, str]: policies: dict[str, str] = {} - for field, enum in ( + for name, enum in ( ("parent_close_policy", ParentClosePolicy), ("cancellation_policy", CancellationPolicy), ): - value = options.get(field) + value = options.get(name) if value is None: continue if isinstance(value, Enum) and not isinstance(value, enum): - raise ValueError(f"child workflow {field} must be a supported policy") + raise ValueError(f"child workflow {name} must be a supported policy") if not isinstance(value, str): - raise ValueError(f"child workflow {field} must be a supported policy") + raise ValueError(f"child workflow {name} must be a supported policy") try: - policies[field] = enum(value).value + policies[name] = enum(value).value except ValueError as error: - raise ValueError(f"child workflow {field} must be a supported policy") from error + raise ValueError(f"child workflow {name} must be a supported policy") from error return policies @@ -109,6 +109,7 @@ class CancellationContext: requested_at: datetime cleanup_deadline_at: datetime lineage: tuple[CancellationLineage, ...] + _replay_clock: Callable[[], datetime] | None = field(default=None, repr=False, compare=False) def __post_init__(self) -> None: object.__setattr__(self, "requester", MappingProxyType(dict(self.requester))) @@ -119,6 +120,19 @@ def deadline(self) -> datetime: """The original immutable cleanup deadline.""" return self.cleanup_deadline_at + def remaining(self) -> float: + """Seconds left at the consumed replay boundary, clamped to zero. + + Available only during the workflow replay that delivered this context. + Detached metadata has no clock. Host time never supplies this value. + """ + if self._replay_clock is None: + raise RuntimeError("cancellation remaining time requires active workflow replay") + return max(0.0, (self.cleanup_deadline_at - self._replay_clock()).total_seconds()) + + def _with_replay_clock(self, clock: Callable[[], datetime]) -> CancellationContext: + return replace(self, _replay_clock=clock) + @classmethod def from_dict(cls, snapshot: Mapping[str, Any]) -> CancellationContext: if snapshot.get("schema") != "durable-workflow.cancellation-context/v1": diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 0bc487b..5191756 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -26,7 +26,9 @@ import random import re import uuid +import weakref from collections.abc import Callable, Generator, Iterable, Mapping, Sequence +from contextvars import ContextVar from copy import copy from dataclasses import dataclass, field from datetime import datetime, timezone @@ -40,6 +42,7 @@ ParentClosePolicy, _canonical_activity_policy, _canonical_child_policies, + _timestamp, ) from .client import WorkflowStreamAppendItem from .errors import ( @@ -1801,6 +1804,9 @@ def run(self, forward: Callable[[Saga], Any]) -> Any: raise +_ACTIVE_WORKFLOW_REPLAY: ContextVar[WorkflowContext | None] = ContextVar("active_workflow_replay", default=None) + + class WorkflowContext: """Replay-safe helper surface passed to workflow ``run`` methods.""" @@ -1822,6 +1828,8 @@ def __init__( self._cancel_requested = bool(cancel_requested) self._cancellation_request_id: str | None = None self._cancellation_context: CancellationContext | None = None + self._cancellation_replay_time: datetime | None = None + self._cancellation_replay_time_available = False self._cancellation_shield_depth = 0 seed = int(hashlib.sha256(run_id.encode()).hexdigest()[:16], 16) self._rng = random.Random(seed) @@ -1926,6 +1934,34 @@ def cancellation_context(self) -> CancellationContext | None: """Original metadata, visible only at committed cancellation delivery.""" return self._cancellation_context + def _observe_cancellation_replay_time(self, event: Mapping[str, Any] | None) -> None: + timestamp = event.get("timestamp") if event is not None else None + if timestamp is None and event is not None: + timestamp = event.get("recorded_at") + self._cancellation_replay_time_available = False + if not isinstance(timestamp, str): + return + try: + recorded_time = _timestamp(timestamp) + except ValueError: + return + if self._cancellation_replay_time is None or recorded_time > self._cancellation_replay_time: + self._cancellation_replay_time = recorded_time + self._cancellation_replay_time_available = True + + def _bind_cancellation_context(self, context: CancellationContext) -> CancellationContext: + reference = weakref.ref(self) + + def replay_time() -> datetime: + active = reference() + if active is None or _ACTIVE_WORKFLOW_REPLAY.get() is not active: + raise RuntimeError("cancellation remaining time requires active workflow replay") + if not active._cancellation_replay_time_available or active._cancellation_replay_time is None: + raise RuntimeError("cancellation replay boundary requires a recorded timestamp") + return active._cancellation_replay_time + + return context._with_replay_clock(replay_time) + @contextlib.contextmanager def cancellation_shield(self) -> Generator[None, None, None]: """Permit cleanup calls without delivering the same request again. @@ -2461,6 +2497,7 @@ class _RecordedStep: shape: str event_types: list[str] details: dict[str, Any] + history_index: int | None = None @dataclass(frozen=True) @@ -3987,15 +4024,18 @@ def _state( # loser complete before or after unrelated successor work without changing # replay binding. selection_resolutions: dict[int, Any] = {} + selection_resolution_indexes: dict[int, int] = {} selection_steps: dict[int, list[_RecordedStep]] = {} selection_payloads: dict[int, list[dict[str, Any]]] = {} selection_opening_payloads: dict[int, dict[str, Any]] = {} selection_condition_terminal_payloads: dict[int, list[dict[str, Any]]] = {} selection_condition_sequences: dict[str, int] = {} selection_markers: list[dict[str, Any]] = [] + selection_marker_indexes: list[int] = [] selection_marker_cursor = 0 consumed_selection_sequences: set[int] = set() cancelled_selection_members: dict[tuple[str, int], dict[str, Any]] = {} + cancelled_selection_indexes: dict[tuple[str, int], int] = {} validated_selection_cancellations: set[tuple[str, int]] = set() authored_selection_handles: dict[int, DurableOperationHandle] = {} @@ -4018,9 +4058,11 @@ def _append_resolved_result(value: Any, shape: str, event: Mapping[str, Any]) -> shape=shape, event_types=event_types, details=details, + history_index=event_index, ) if _selection_group_metadata(payload) is not None: selection_resolutions[workflow_sequence] = value + selection_resolution_indexes[workflow_sequence] = event_index _append_selection_step(shape, event, fallback_sequence=workflow_sequence) return resolved_results.append(value) @@ -4059,6 +4101,7 @@ def _recorded_step( shape=shape, event_types=event_types, details=details, + history_index=event_index, ) def _append_selection_step( @@ -4120,6 +4163,7 @@ def _append_selection_condition_resolution( fallback_sequence=selection_sequence, ) selection_resolutions[selection_sequence] = value + selection_resolution_indexes[selection_sequence] = event_index return True def _assert_step_matches(command: Any, step: _RecordedStep) -> None: @@ -4500,6 +4544,7 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: elif etype == "SelectionResolved": if isinstance(payload.get("selection_group_id"), str): selection_markers.append(dict(payload)) + selection_marker_indexes.append(event_index) elif etype == "SelectionOperationCancelled": group_id = payload.get("selection_group_id") member_base = _parallel_group_integer(payload, "member_base_sequence") @@ -4515,6 +4560,7 @@ def _receiver_condition_wait_bindings() -> dict[int, str | None]: detail="conflicting cancellation markers target the same durable member", ) cancelled_selection_members[cancellation_key] = cancellation_payload + cancelled_selection_indexes.setdefault(cancellation_key, event_index) elif etype in ("SideEffectRecorded", "ChildRunCompleted"): shape = "side effect" if etype == "SideEffectRecorded" else "child workflow" _append_resolved_result( @@ -5508,7 +5554,12 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl cancellation_consumed = True ctx._cancel_requested = True ctx._cancellation_request_id = boundary.request_id - ctx._cancellation_context = cancellation.request.context if cancellation.request is not None else None + metadata = cancellation.request.context if cancellation.request is not None else None + if metadata is not None: + ctx._observe_cancellation_replay_time( + events[cancellation.delivery_index] if cancellation.delivery_index is not None else None, + ) + ctx._cancellation_context = ctx._bind_cancellation_context(metadata) _apply_due_receivers() if isinstance(command, DurableOperationHandle): authored_sequence += 1 @@ -5521,6 +5572,26 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl except StopIteration as stop: return _terminal_state(stop.value, include_pending=True) + def _advance_cancellation_clock(indexes: Iterable[int | None]) -> None: + if ctx._cancellation_context is None: + return + for index in sorted({index for index in indexes if index is not None}): + ctx._observe_cancellation_replay_time(events[index]) + + def _advance_selection_clock(base: int, size: int, failure: BaseException | None) -> None: + indexes = [ + selection_resolution_indexes[sequence] for sequence in range(base, base + size) + if sequence in selection_resolution_indexes + ] + if failure is not None: + failure_index = next( + selection_resolution_indexes[sequence] for sequence in range(base, base + size) + if selection_resolutions.get(sequence) is failure + ) + indexes = [index for index in indexes if index <= failure_index] + _advance_cancellation_clock(indexes) + + replay_token = _ACTIVE_WORKFLOW_REPLAY.set(ctx) try: while True: # Cursor-0 receivers are start-boundary events. Enter run() once @@ -5665,6 +5736,7 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl ["SelectionResolved"], detail="the committed winner has no terminal member history", ) + _advance_cancellation_clock([selection_marker_indexes[selection_marker_cursor - 1]]) for sequence in range( winner.base_sequence, winner.base_sequence + winner.size, @@ -5709,8 +5781,20 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl for offset, child_command in enumerate(cmd): _assert_next_step_matches(child_command, offset) vals = resolved_results[result_cursor : result_cursor + needed] + consumed_steps = recorded_steps[result_cursor : result_cursor + needed] result_cursor += needed failed = _first_yield_failure(vals) + consumed_indexes = [step.history_index for step in consumed_steps] + if failed is not None: + failure_index = next( + step.history_index for step, value in zip(consumed_steps, vals, strict=True) + if value is failed + ) + consumed_indexes = [ + index for index in consumed_indexes + if index is not None and failure_index is not None and index <= failure_index + ] + _advance_cancellation_clock(consumed_indexes) if failed is not None: try: advanced_cmd = gen.throw(failed) @@ -5727,6 +5811,9 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl if isinstance(cmd, DurableOperationHandle): _assert_authored_selection_handle(cmd) if _selection_cancellation_for_handle(cmd) is not None: + _advance_cancellation_clock([ + cancelled_selection_indexes[(cmd.selection_group_id, cmd.base_sequence)], + ]) try: advanced_cmd = gen.throw( DurableOperationCancelled( @@ -5757,6 +5844,7 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl if not resolved: ctx.logger._set_replaying(False) return _state(pending) + _advance_selection_clock(cmd.base_sequence, cmd.size, failure) if failure is not None: try: advanced_cmd = gen.throw(failure) @@ -5939,6 +6027,8 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl if isinstance(opened_id, str): resolution = wait_resolutions.get(opened_id) _apply_condition_wait_receivers(opened_id) + if resolution is not None: + _advance_cancellation_clock([condition_wait_terminal_indexes.get(opened_id)]) next_wait_index = wait_yield_count + 1 has_reopened_same_wait = ( opened is not None @@ -6000,6 +6090,8 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl if result_cursor < len(resolved_results): _assert_next_step_matches(cmd) val = resolved_results[result_cursor] + if not isinstance(cmd, RecordLocalActivity) or prepare_local_activities: + _advance_cancellation_clock([recorded_steps[result_cursor].history_index]) result_cursor += 1 if isinstance(val, ActivityFailed | ChildWorkflowFailed): try: @@ -6044,3 +6136,5 @@ def _consume_cancellation(command: Any, boundary: CancellationDelivery) -> _Repl return _state(pending + [_fail_workflow_from_exception(exc)]) except Exception as exc: return _state(pending + [_fail_workflow_from_exception(exc)]) + finally: + _ACTIVE_WORKFLOW_REPLAY.reset(replay_token) diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index 3a76f26..c0ce04d 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -50,8 +50,16 @@ def run(self, ctx: Any, kind: str, cancellation_policy: str | None = None) -> An else: yield ctx.start_timer(300) except WorkflowCancelled as error: + if kind == "clock_timer": + assert error.context is ctx.cancellation_context and error.context is not None + print(json.dumps({"phase": "remaining-delivery", "remaining": error.context.remaining(), + "context": error.context.to_dict()}), flush=True) with ctx.cancellation_shield(): yield ctx.local_activity("tests.python-cooperative-cleanup", [error.request_id]) + if kind == "clock_timer": + assert error.context is not None + print(json.dumps({"phase": "remaining-cleanup", "remaining": error.context.remaining(), + "context": error.context.to_dict()}), flush=True) return error.request_id return "not cancelled" @@ -766,13 +774,15 @@ async def test_killed_process_reclaims_cleanup_in_a_new_process( await seed._register() try: handle = await client.start_workflow( - workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, input=["timer"], + workflow_type="tests.python-cooperative-cleanup", workflow_id=queue, task_queue=queue, + input=["clock_timer"], ) accepted = await handle.request_cancellation(cleanup_timeout_seconds=120) original = accepted["cancellation_request"] first = await native_process(queue, f"{queue}-killed", "hold") processes.append(first) claim = await process_event(first, "claim") + original_remaining = await process_event(first, "remaining-delivery") cleanup = await process_event(first, "cleanup") assert cleanup["request_id"] == original["request_id"] before = await events(handle) @@ -786,11 +796,26 @@ async def test_killed_process_reclaims_cleanup_in_a_new_process( assert reclaimed["worker_id"] != claim["worker_id"] assert reclaimed["request_id"] == original["request_id"] assert reclaimed["cleanup_deadline_at"] == original["cleanup_deadline_at"] + replacement_remaining = await process_event(replacement, "remaining-delivery") + assert replacement_remaining == original_remaining resumed = await process_event(replacement, "cleanup") assert resumed["request_id"] == original["request_id"] + completed_remaining = await process_event(replacement, "remaining-cleanup") + assert completed_remaining["context"] == original_remaining["context"] + assert 0 < completed_remaining["remaining"] < original_remaining["remaining"] assert (await process_event(replacement, "finished"))["committed"] is True assert await asyncio.wait_for(replacement.wait(), timeout=10) == 0 - await assert_cancelled_cleanup(handle, original["request_id"]) + history = await assert_cancelled_cleanup(handle, original["request_id"]) + deadline = datetime.fromisoformat(original["cleanup_deadline_at"].replace("Z", "+00:00")) + delivered = next(event for event in history if event["event_type"] == "CooperativeCancellationDelivered") + completed = next(event for event in history if event["event_type"] == "ActivityCompleted") + for event, observation in ((delivered, original_remaining), (completed, completed_remaining)): + recorded = datetime.fromisoformat(event["timestamp"].replace("Z", "+00:00")) + assert observation["remaining"] == pytest.approx((deadline - recorded).total_seconds(), abs=1e-6) + remaining_observations = { + "original": original_remaining, "replacement": replacement_remaining, "completed": completed_remaining, + } + print(f"SIGKILL replay remaining-time observations: {json.dumps(remaining_observations)}") finally: for process in processes: if process.returncode is None: diff --git a/tests/test_cancellation_remaining_time.py b/tests/test_cancellation_remaining_time.py new file mode 100644 index 0000000..db7d5c1 --- /dev/null +++ b/tests/test_cancellation_remaining_time.py @@ -0,0 +1,265 @@ +from __future__ import annotations + +from copy import deepcopy +from datetime import datetime, timezone +from typing import Any + +import pytest + +from durable_workflow import CancellationContext, serializer +from durable_workflow.errors import ActivityFailed, DurableOperationCancelled, WorkflowCancelled +from durable_workflow.workflow import ( + CompleteWorkflow, + RecordLocalActivity, + RecordSideEffect, + ScheduleActivity, + WorkflowContext, + replay, +) +from tests.test_cancellation_context import delivery, request, snapshot +from tests.test_cooperative_cancellation import completed_activity +from tests.test_durable_selection import _activity_completed, _activity_scheduled, _winner_marker + + +def history() -> list[dict[str, Any]]: + return [ + {"event_type": "WorkflowStarted", "payload": {"timestamp": "2025-01-01T00:00:00Z"}}, + request(), + {**delivery(), "timestamp": "2026-10-01T00:00:08Z"}, + ] + + +def test_detached_metadata_has_no_remaining_time_clock() -> None: + context = CancellationContext.from_dict(snapshot()) + with pytest.raises(RuntimeError, match="active workflow replay"): + context.remaining() + + +def test_first_execution_and_cold_replay_keep_the_same_budget_through_synchronous_results() -> None: + observations: list[tuple[float, float]] = [] + contexts: list[CancellationContext] = [] + calls: list[str] = [] + + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is ctx.cancellation_context + assert error.context is not None + contexts.append(error.context) + before = error.context.remaining() + with ctx.cancellation_shield(): + yield ctx.side_effect(lambda: calls.append("once") or "recorded") + observations.append((before, error.context.remaining())) + yield ctx.schedule_activity("cleanup", []) + assert ctx.now() == datetime(2025, 1, 1, tzinfo=timezone.utc) + return {"remaining": error.context.remaining(), "context": error.context.to_dict()} + + events = history() + first = replay(Cleanup, events, [], run_id="run-1") + assert isinstance(first.commands[0], RecordSideEffect) + assert isinstance(first.commands[1], ScheduleActivity) + events.extend([ + {"event_type": "SideEffectRecorded", "timestamp": "2026-10-01T00:00:28Z", "payload": { + "sequence": 2, "result": serializer.envelope("recorded"), + }}, + {**completed_activity(3, "cleanup", "cleaned"), "timestamp": "2026-10-01T00:00:25Z"}, + {"event_type": "WorkflowTaskCompleted", "timestamp": "2026-10-01T00:00:29Z", "payload": {}}, + ]) + for _restart in range(2): + assert replay(Cleanup, events, [], run_id="run-1").commands == [ + CompleteWorkflow({"remaining": 5.123456, "context": snapshot()}), + ] + assert observations == [(22.123456, 22.123456)] * 3 + assert calls == ["once"] + for context in contexts: + with pytest.raises(RuntimeError, match="active workflow replay"): + context.remaining() + + +def test_completed_parallel_group_uses_consumed_members_without_regressing_on_clock_skew() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + yield [ctx.schedule_activity("one", []), ctx.schedule_activity("two", [])] + return error.context.remaining() + + events = history() + [ + {**completed_activity(2, "one", "one"), "timestamp": "2026-10-01T00:00:25Z"}, + {**completed_activity(3, "two", "two"), "timestamp": "2026-10-01T00:00:20Z"}, + ] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(5.123456)] + + +def test_parallel_failure_does_not_consume_a_future_sibling_completion() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + try: + yield [ctx.schedule_activity("one", []), ctx.schedule_activity("two", [])] + except ActivityFailed: + return error.context.remaining() + + events = history() + [ + {"event_type": "ActivityFailed", "timestamp": "2026-10-01T00:00:12Z", "payload": { + "sequence": 2, "activity_type": "one", "message": "failed", "exception_class": "RuntimeError", + }}, + {**completed_activity(3, "two", "two"), "timestamp": "2026-10-01T00:00:29Z"}, + ] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(18.123456)] + + +@pytest.mark.parametrize("timestamp", [None, "invalid", "2026-10-01T00:00:08", "2026-02-30T00:00:08Z"]) +def test_missing_or_invalid_delivery_time_refuses_a_host_clock(timestamp: str | None) -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with pytest.raises(RuntimeError, match="recorded timestamp"): + error.context.remaining() + return "refused" + + events = history() + events[-1]["timestamp"] = timestamp + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow("refused")] + + +def test_recorded_expiry_clamps_remaining_to_zero_and_accepts_timezone_offsets() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + yield ctx.start_timer(10) + return error.context.remaining() + + events = history() + [{"event_type": "TimerFired", "recorded_at": "2026-09-30T20:00:31-04:00", "payload": { + "sequence": 2, "timer_kind": "durable_timer", "duration_ms": 10000, + }}] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(0.0)] + + +def selection_event(event: dict[str, Any], timestamp: str) -> dict[str, Any]: + shifted = deepcopy(event) + shifted["timestamp"] = f"2026-10-01T00:00:{timestamp}Z" + payload = shifted["payload"] + for item in [payload, *payload.get("parallel_group_path", [])]: + for key in ("sequence", "parallel_group_base_sequence", "selection_member_base_sequence", + "selection_group_base_sequence", "member_base_sequence"): + if key in item: + item[key] += 1 + for key in ("parallel_group_id", "selection_group_id"): + if key in item: + item[key] = "select-calls:2:2" + return shifted + + +def selection_history() -> list[dict[str, Any]]: + return history() + [ + selection_event(_activity_scheduled(0, "slow"), "09"), + selection_event(_activity_scheduled(1, "fast"), "09"), + selection_event(_activity_completed(1, "fast", "fast"), "10"), + selection_event(_winner_marker(), "12"), + ] + + +def test_selection_advances_at_the_winner_and_awaited_handle_without_future_loser_lookahead() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + selected = yield ctx.select({ + "slow": ctx.schedule_activity("slow-activity", []), + "fast": ctx.schedule_activity("fast-activity", []), + }) + at_winner = error.context.remaining() + yield selected.winner.await_result() + at_old_result = error.context.remaining() + yield selected.handles["slow"].await_result() + return [at_winner, at_old_result, error.context.remaining()] + + events = selection_history() + [selection_event(_activity_completed(0, "slow", "slow"), "20")] + assert replay(Cleanup, events, [], run_id="run-1").commands == [ + CompleteWorkflow([18.123456, 18.123456, 10.123456]), + ] + + +def test_cancelled_handle_uses_the_first_durable_receipt_despite_identical_redelivery() -> None: + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + selected = yield ctx.select({ + "slow": ctx.schedule_activity("slow-activity", []), + "fast": ctx.schedule_activity("fast-activity", []), + }) + yield selected.handles["slow"].cancel() + try: + yield selected.handles["slow"].await_result() + except DurableOperationCancelled: + return error.context.remaining() + + cancelled = {"event_type": "SelectionOperationCancelled", "timestamp": "2026-10-01T00:00:15Z", + "payload": {"selection_group_id": "select-calls:2:2", "member_key": "slow", "member_index": 0, + "member_base_sequence": 2, "member_size": 1, "operation_kind": "activity", + "operation_identity": "activity-slow"}} + events = selection_history() + [cancelled, {**cancelled, "timestamp": "2026-10-01T00:00:28Z"}] + assert replay(Cleanup, events, [], run_id="run-1").commands == [CompleteWorkflow(15.123456)] + + +def test_inline_local_result_persistence_preserves_the_first_execution_budget() -> None: + observations: list[float] = [] + calls: list[str] = [] + + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + try: + yield ctx.start_timer(60) + except WorkflowCancelled as error: + assert error.context is not None + with ctx.cancellation_shield(): + yield ctx.local_activity("inline", []) + observations.append(error.context.remaining()) + yield ctx.schedule_activity("cleanup", []) + return error.context.remaining() + + def execute(command: RecordLocalActivity) -> str: + calls.append("once") + command.arguments_envelope = serializer.envelope(command.arguments) + command.result_envelope = serializer.envelope("local") + command.outcome = {"outcome": "completed", "attempts": []} + return "local" + + events = history() + first = replay(Cleanup, events, [], run_id="run-1", local_activity_executor=execute) + assert isinstance(first.commands[0], RecordLocalActivity) + events.extend([ + {"event_type": "ActivityCompleted", "timestamp": "2026-10-01T00:00:28Z", "payload": { + "sequence": 2, "activity_type": "inline", "execution_mode": "local", "result": serializer.envelope("local"), + }}, + {**completed_activity(3, "cleanup", "cleaned"), "timestamp": "2026-10-01T00:00:25Z"}, + ]) + assert replay(Cleanup, events, [], run_id="run-1", local_activity_executor=execute).commands == [ + CompleteWorkflow(5.123456), + ] + assert observations == [22.123456, 22.123456] + assert calls == ["once"] From 0e1c427eae2f6120155a5cf3c9c1ae79ba7a5756 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Fri, 2 Oct 2026 21:36:52 +0000 Subject: [PATCH 33/41] Qualify remaining budgets across prepared cleanup recovery --- docs/cooperative-cancellation-design.md | 8 +++-- .../test_cooperative_cancellation.py | 8 ++--- .../test_prepared_local_activity.py | 29 ++++++++++++++++++- 3 files changed, 36 insertions(+), 9 deletions(-) diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index db5aee9..df04fd6 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -233,9 +233,11 @@ or invalid boundary timestamps also fail, without a host-time fallback. The runtime supervisor independently enforces the actual deadline and task ownership even when authoring code cannot run. -The connected process-loss scenario records remaining time before SIGKILL, -requires the same value in the replacement worker, then checks the value after -prepared cleanup against its committed history timestamp and original deadline. +Connected process-loss scenarios record remaining time before SIGKILL and +require the same value and metadata in the replacement worker. Legacy inline +cleanup preserves that value through result persistence. Sequential and atomic +prepared cleanup, each under an original 30-second deadline, check the final +value against the committed completion timestamp and original deadline. Rust helpers, explicit local operation policies, nested scopes and competitive qualification still need completion. diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index c0ce04d..f92f2cf 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -802,16 +802,14 @@ async def test_killed_process_reclaims_cleanup_in_a_new_process( assert resumed["request_id"] == original["request_id"] completed_remaining = await process_event(replacement, "remaining-cleanup") assert completed_remaining["context"] == original_remaining["context"] - assert 0 < completed_remaining["remaining"] < original_remaining["remaining"] + assert completed_remaining["remaining"] == original_remaining["remaining"] > 0 assert (await process_event(replacement, "finished"))["committed"] is True assert await asyncio.wait_for(replacement.wait(), timeout=10) == 0 history = await assert_cancelled_cleanup(handle, original["request_id"]) deadline = datetime.fromisoformat(original["cleanup_deadline_at"].replace("Z", "+00:00")) delivered = next(event for event in history if event["event_type"] == "CooperativeCancellationDelivered") - completed = next(event for event in history if event["event_type"] == "ActivityCompleted") - for event, observation in ((delivered, original_remaining), (completed, completed_remaining)): - recorded = datetime.fromisoformat(event["timestamp"].replace("Z", "+00:00")) - assert observation["remaining"] == pytest.approx((deadline - recorded).total_seconds(), abs=1e-6) + recorded = datetime.fromisoformat(delivered["timestamp"].replace("Z", "+00:00")) + assert original_remaining["remaining"] == pytest.approx((deadline - recorded).total_seconds(), abs=1e-6) remaining_observations = { "original": original_remaining, "replacement": replacement_remaining, "completed": completed_remaining, } diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index 976c22c..04baaea 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -62,6 +62,11 @@ def run(self, ctx: Any, marker: str, group: bool = False) -> Any: else: yield ctx.local_activity("tests.python-prepared-blocked", [marker, "work", None]) except WorkflowCancelled as error: + assert error.context is ctx.cancellation_context and error.context is not None + delivery_path = Path(marker + ".remaining-delivery-" + str(os.getpid()) + ".json") + pending = delivery_path.with_suffix(".writing") + pending.write_text(json.dumps({"context": error.context.to_dict(), "remaining": error.context.remaining()})) + pending.replace(delivery_path) with ctx.cancellation_shield(): if group: yield [ctx.local_activity( @@ -73,6 +78,10 @@ def run(self, ctx: Any, marker: str, group: bool = False) -> Any: "tests.python-prepared-blocked", [marker, "cleanup", error.request_id], retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, ) + final_path = Path(marker + ".remaining-final.json") + pending = final_path.with_suffix(".writing") + pending.write_text(json.dumps({"context": error.context.to_dict(), "remaining": error.context.remaining()})) + pending.replace(final_path) return error.request_id return "not cancelled" @@ -256,6 +265,8 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: cleanup = [await remote_marker(Path(marker + "." + phase + "." + queue + "-first")) for phase in cleanup_phases] assert all(item["request_id"] == original["request_id"] for item in cleanup) + original_remaining = await remote_marker(Path(marker + ".remaining-delivery-" + str(first.pid) + ".json")) + assert original_remaining["context"] == original_context for item in work: await callback_gone(item["callback_pid"]) first.kill() @@ -265,13 +276,26 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: duplicate = await handle.request_cancellation(cleanup_timeout_seconds=90) for field in ("request_id", "requested_at", "cleanup_deadline_at"): assert duplicate["cancellation_request"][field] == original[field] - await owner(queue + "-replacement", "finish") + successor = await owner(queue + "-replacement", "finish") deadline = datetime.fromisoformat(original["cleanup_deadline_at"].replace("Z", "+00:00")) timeout = (deadline - datetime.now(deadline.tzinfo)).total_seconds() assert timeout > 0 with pytest.raises(WorkflowCancelled): await handle.result(timeout=timeout) history = await events(handle) + replacement_remaining = await remote_marker( + Path(marker + ".remaining-delivery-" + str(successor.pid) + ".json"), + ) + completed_remaining = await remote_marker(Path(marker + ".remaining-final.json")) + assert replacement_remaining == original_remaining + assert completed_remaining["context"] == original_context + assert 0 < completed_remaining["remaining"] < original_remaining["remaining"] + delivered = next(event for event in history if event["event_type"] == "CooperativeCancellationDelivered") + completed = max(datetime.fromisoformat(event["timestamp"].replace("Z", "+00:00")) + for event in history if event["event_type"] == "ActivityCompleted") + delivered_at = datetime.fromisoformat(delivered["timestamp"].replace("Z", "+00:00")) + assert original_remaining["remaining"] == pytest.approx((deadline - delivered_at).total_seconds(), abs=1e-6) + assert completed_remaining["remaining"] == pytest.approx((deadline - completed).total_seconds(), abs=1e-6) kinds = [event["event_type"] for event in history] assert kinds.count("CooperativeCancellationDelivered") == kinds.count("WorkflowCancelled") == 1 assert kinds.count("ActivityCancellationAcknowledged") == len(work) @@ -299,6 +323,9 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: print("prepared physically joined callbacks: " + json.dumps({ "work": work, "killed_cleanup": cleanup, "replacement_cleanup": replacement, })) + print("prepared remaining-time observations: " + json.dumps({ + "original": original_remaining, "replacement": replacement_remaining, "completed": completed_remaining, + })) print("prepared cleanup recovery receipts: " + json.dumps(receipts)) print("prepared cleanup recovery history: " + json.dumps(history)) finally: From afa7cfaabd0d7fad8049f7b6e3b4c19ce15d0aa2 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sat, 3 Oct 2026 01:01:40 +0000 Subject: [PATCH 34/41] Support prepared local Activity cancellation policies --- docs/cooperative-cancellation-design.md | 30 ++- src/durable_workflow/worker.py | 39 ++++ src/durable_workflow/workflow.py | 64 +++++- .../test_prepared_local_activity.py | 30 ++- .../test_prepared_local_activity_policies.py | 206 ++++++++++++++++++ tests/test_prepared_local_activity_worker.py | 27 ++- 6 files changed, 380 insertions(+), 16 deletions(-) create mode 100644 tests/test_prepared_local_activity_policies.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index df04fd6..103e6d9 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -92,9 +92,8 @@ the independent work. It does not extend the parent's cleanup budget. Every explicit remote policy requires negotiated cooperation, protocol 1.20 and a compatible installed backend. Missing worker capability is diagnosed -before submission. Server also checks the original immutable claim. Local -Activity policy authoring remains unavailable while its lifetime and admission -contract are unfinished. +before submission. Server also checks the original immutable claim. Prepared +local Activity policies use their separate admission contract below. Replay compares authored policy with original canonical history for ordinary, parallel and selection calls and cancellation delivery. Later events that omit @@ -153,6 +152,31 @@ effects already performed. ## Prepared sequential local callbacks +`ctx.local_activity()` accepts `cancellation_policy=CancellationPolicy.TRY_CANCEL` +or `CancellationPolicy.WAIT_CANCELLATION_COMPLETED`. Omission, including Python's +optional `None`, preserves historical Try behavior and omits the wire field. +Explicit policies require prepared execution and Server discovery of +`prepared_local_activity_cancellation_policies`. The Worker advertises its policy +consumer only for the discovered installed policies. Local `ABANDON` is refused +because a callback owned by this workflow worker cannot outlive that ownership +under the prepared contract. + +```python +yield ctx.local_activity( + "release-reservation", [], + cancellation_policy=CancellationPolicy.WAIT_CANCELLATION_COMPLETED, +) +``` + +Wait parks cancellation delivery until the original local attempt's stop receipt +is recorded. Both supported policies physically stop and join the owned callback +before acknowledgment, without requiring application heartbeats. Neither grants +a new cleanup budget. Replay compares the original policy for completed and +unresolved calls, groups and the committed delivery boundary. Unsupported +policies refuse the entire authored group before callbacks or checkpoint +submission. Queries, updates and validators use the same negotiated replay +consumer. Explicit policies are unavailable through the legacy inline path. + The source candidate can explicitly request both `cooperative_cancellation` and `prepared_local_activities` in Worker capabilities. Registration requires source protocol 1.20 and actual Server discovery of its installed admission bridge. diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 19b4d5b..6029d9d 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -1055,6 +1055,11 @@ def __init__( self._cooperative_cancellation_supported = False self._prepared_local_activities_supported = False self._prepared_local_activity_groups_supported = False + self._local_activity_cancellation_policies: tuple[str, ...] = () + if "prepared_local_activity_cancellation_policies" in self.capabilities and ( + "prepared_local_activities" not in self.capabilities + ): + raise ValueError("prepared local cancellation policies require prepared_local_activities capability") if "prepared_local_activity_groups" in self.capabilities and ( "prepared_local_activities" not in self.capabilities ): @@ -1255,6 +1260,21 @@ async def _register(self) -> None: ) if "prepared_local_activity_groups" in self.capabilities and not self._prepared_local_activity_groups_supported: raise RuntimeError("prepared_local_group_not_supported: Server must advertise its installed atomic bridge") + policies = ( + server_capabilities.get("prepared_local_activity_cancellation_policies") + if isinstance(server_capabilities, Mapping) else None + ) + self._local_activity_cancellation_policies = tuple( + policy for policy in ("try_cancel", "wait_cancellation_completed") + if isinstance(policies, list) and policy in policies + ) if self._prepared_local_activities_supported else () + if "prepared_local_activity_cancellation_policies" in self.capabilities and ( + not self._local_activity_cancellation_policies + ): + raise RuntimeError( + "prepared_local_activity_cancellation_policy_not_supported: " + "Server must advertise installed prepared policies", + ) self._validate_cooperative_activity_handlers() self._query_tasks_supported = _server_supports_query_tasks(info) self._workflow_memo_updates_supported = _server_supports_workflow_memo_updates(info) @@ -1289,6 +1309,11 @@ async def _register(self) -> None: capabilities.append(WORKFLOW_UPDATES_CAPABILITY) capabilities.append(MESSAGE_STREAMS_CAPABILITY) capabilities.extend(self.capabilities) + if ( + self._local_activity_cancellation_policies + and "prepared_local_activity_cancellation_policies" not in capabilities + ): + capabilities.append("prepared_local_activity_cancellation_policies") if PORTABLE_WORKER_AFFINITY_CAPABILITY_MANIFEST["worker_sessions"]["supported"]: capabilities.append("worker_sessions") @@ -1314,6 +1339,10 @@ async def _register(self) -> None: "supported": True, "minimum_protocol_version": "1.20", "implementation": "durable_atomic_all_admission", }} if self._prepared_local_activity_groups_supported else {}), + **({"prepared_local_activity_cancellation_policies": { + "supported": True, "minimum_protocol_version": "1.20", + "implementation": "prepared_local_policy_admission_and_replay", + }} if self._local_activity_cancellation_policies else {}), }, task_slots=self._current_task_slots(), process_metrics=self._current_process_metrics(), @@ -1568,6 +1597,7 @@ async def _replay_workflow_claim( cancellation_request=task.get("cancellation_request"), local_activity_executor=execute_local, prepare_local_activities=self._prepared_local_activities_supported, prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, ) if outcome.prepared_local_activity_group is not None: history = await self._execute_prepared_local_activity_group(task, history, outcome) @@ -2326,6 +2356,9 @@ async def _run_workflow_task_core(self, task: dict[str, Any]) -> list[dict[str, payload_codec=codec, external_storage=self.external_storage, external_storage_cache=self.external_storage_cache, + prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, ) command = update_command.to_server_command( self.task_queue, @@ -3146,6 +3179,9 @@ async def _run_query_task_core(self, task: dict[str, Any], *, client: Client | N payload_codec=codec, external_storage=self.external_storage, external_storage_cache=self.external_storage_cache, + prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, ) if inspect.isawaitable(result): result = await result @@ -3637,6 +3673,9 @@ async def _run_update_validation_task(self, task: dict[str, Any]) -> str: payload_codec=codec, external_storage=self.external_storage, external_storage_cache=self.external_storage_cache, + prepare_local_activities=self._prepared_local_activities_supported, + prepare_local_activity_groups=self._prepared_local_activity_groups_supported, + local_activity_cancellation_policies=self._local_activity_cancellation_policies, ) if inspect.isawaitable(result): result = await result diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 5191756..ccf127c 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -594,12 +594,19 @@ class RecordLocalActivity: start_to_close_timeout: int | None = None schedule_to_close_timeout: int | None = None heartbeat_timeout: int | None = None + cancellation_policy: str | CancellationPolicy | None = None outcome: dict[str, Any] | None = field(default=None, init=False, repr=False) arguments_envelope: dict[str, Any] | None = field(default=None, init=False, repr=False) result_envelope: dict[str, Any] | None = field(default=None, init=False, repr=False) _parallel_group_path: list[dict[str, Any]] | None = field(default=None, init=False, repr=False, compare=False) def __post_init__(self) -> None: + try: + self.cancellation_policy = _canonical_activity_policy(self.cancellation_policy) + except ValueError as error: + raise ValueError("local activity cancellation_policy must be a supported prepared policy") from error + if self.cancellation_policy == CancellationPolicy.ABANDON.value: + raise ValueError("local activity abandon is not supported by prepared callback ownership") if not isinstance(self.activity_type, str): raise TypeError("local activity type must be a string") self.activity_type = self.activity_type.strip() @@ -633,6 +640,11 @@ def to_server_command( size_warning: serializer.PayloadSizeWarningConfig | None = serializer.DEFAULT_PAYLOAD_SIZE_WARNING, warning_context: PayloadWarningContext = None, ) -> dict[str, Any]: + if self.cancellation_policy is not None: + raise LocalActivityExecutionAborted( + "prepared_local_activity_cancellation_policy_not_supported: " + "explicit policies require prepared admission", + ) if self.outcome is None or self.arguments_envelope is None: raise ValueError("local activity has no recorded terminal outcome") status = self.outcome.get("outcome") @@ -675,6 +687,8 @@ def descriptor(self, payload_codec: str) -> dict[str, Any]: } if command.retry_policy is not None: descriptor["retry_policy"] = dict(command.retry_policy) + if command.cancellation_policy is not None: + descriptor["cancellation_policy"] = command.cancellation_policy for name in ("start_to_close_timeout", "schedule_to_close_timeout", "heartbeat_timeout"): value = getattr(command, name) if value is not None: @@ -2015,6 +2029,7 @@ def local_activity( start_to_close_timeout: int | None = None, schedule_to_close_timeout: int | None = None, heartbeat_timeout: int | None = None, + cancellation_policy: str | CancellationPolicy | None = None, ) -> RecordLocalActivity: """Yield an activity executed in this workflow worker process.""" return RecordLocalActivity( @@ -2024,6 +2039,7 @@ def local_activity( start_to_close_timeout=start_to_close_timeout, schedule_to_close_timeout=schedule_to_close_timeout, heartbeat_timeout=heartbeat_timeout, + cancellation_policy=cancellation_policy, ) def start_timer(self, seconds: int) -> StartTimer: @@ -2777,6 +2793,7 @@ def replay( local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, prepare_local_activities: bool = False, prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> ReplayOutcome: return _replay_state( workflow_cls, @@ -2793,6 +2810,7 @@ def replay( local_activity_executor=local_activity_executor, prepare_local_activities=prepare_local_activities, prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ).outcome @@ -2808,6 +2826,9 @@ def query_state( payload_codec: str | None = None, external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> Any: """Replay a workflow to current state and invoke a registered query. @@ -2830,6 +2851,9 @@ def query_state( external_storage=external_storage, external_storage_cache=external_storage_cache, stop_at_uncommitted_cancellation=True, + prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ) except Exception as exc: raise QueryFailed(f"workflow replay failed before query: {exc}") from exc @@ -2865,6 +2889,9 @@ def apply_update( payload_codec: str | None = None, external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> CompleteUpdate | FailUpdate: """Replay current workflow state and run one accepted update handler. @@ -2885,6 +2912,9 @@ def apply_update( payload_codec=payload_codec, external_storage=external_storage, external_storage_cache=external_storage_cache, + prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ) except Exception as exc: return _fail_update_from_exception( @@ -2966,6 +2996,9 @@ def validate_update( payload_codec: str | None = None, external_storage: ExternalStorageDriver | None = None, external_storage_cache: ExternalPayloadCache | None = None, + prepare_local_activities: bool = False, + prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), ) -> Any: """Replay state and invoke only the declared pre-accept update validator. @@ -2986,6 +3019,9 @@ def validate_update( payload_codec=payload_codec, external_storage=external_storage, external_storage_cache=external_storage_cache, + prepare_local_activities=prepare_local_activities, + prepare_local_activity_groups=prepare_local_activity_groups, + local_activity_cancellation_policies=local_activity_cancellation_policies, ) except Exception as exc: raise UpdateValidationFailed( @@ -3625,6 +3661,13 @@ def _recorded_detail_mismatch(command: Any, step: _RecordedStep) -> str | None: f"Recorded local activity_type {recorded!r}, but current workflow " f"requested {command.activity_type!r}." ) + actual_policy = command.cancellation_policy or "try_cancel" + recorded_policy = step.details.get("cancellation_policy", "try_cancel") + if recorded_policy != actual_policy: + return ( + f"local_activity_cancellation_policy_changed: recorded {recorded_policy!r}, " + f"but current workflow requested {actual_policy!r}." + ) elif isinstance(command, StartChildWorkflow): recorded = step.details.get("workflow_type") or step.details.get("child_workflow_type") if isinstance(recorded, str) and recorded != command.workflow_type: @@ -3801,6 +3844,7 @@ def _replay_state( local_activity_executor: Callable[[RecordLocalActivity], Any] | None = None, prepare_local_activities: bool = False, prepare_local_activity_groups: bool = False, + local_activity_cancellation_policies: tuple[str, ...] = (), stop_at_uncommitted_cancellation: bool = False, ) -> _ReplayState: if prepare_local_activity_groups and not prepare_local_activities: @@ -5407,6 +5451,23 @@ def _contains_local_activity(operation: Any) -> bool: return any(_contains_local_activity(member) for _, member in operation.operations) return False + def _validate_local_activity_policies(operation: Any) -> None: + if isinstance(operation, RecordLocalActivity) and operation.cancellation_policy is not None: + if ( + not prepare_local_activities + or operation.cancellation_policy not in local_activity_cancellation_policies + ): + raise LocalActivityExecutionAborted( + "prepared_local_activity_cancellation_policy_not_supported: " + f"{operation.cancellation_policy} requires installed Server discovery and a prepared consumer", + ) + elif isinstance(operation, list): + for member in operation: + _validate_local_activity_policies(member) + elif isinstance(operation, SelectGroup): + for _, member in operation.operations: + _validate_local_activity_policies(member) + def _prepared_call(command: RecordLocalActivity, sequence: int) -> PreparedLocalActivityCall: cleanup: dict[str, str] | None = None if cancellation.request is not None: @@ -5461,7 +5522,7 @@ def _prepared_group(commands: list[Any]) -> PreparedLocalActivityGroup | None: continue # A complete path is admission authority. Legacy metadata-poor # history cannot authorize fresh Source callbacks. - details = _recorded_step_details(payload) + details = {**details_by_sequence.get(sequence, {}), **_recorded_step_details(payload)} if payload.get("parallel_group_path") != command._parallel_group_path: raise NonDeterministicReplayError( sequence, "complete authored parallel group path", [kind], @@ -5610,6 +5671,7 @@ def _advance_selection_clock(base: int, size: int, failure: BaseException | None return _terminal_state(stop.value, include_pending=True) first = False _apply_due_receivers() + _validate_local_activity_policies(cmd) if prepare_local_activities and isinstance(cmd, list | SelectGroup) and _contains_local_activity(cmd) and ( not prepare_local_activity_groups or isinstance(cmd, SelectGroup) ): diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index 04baaea..b5d78a9 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -29,6 +29,9 @@ async def prepared_runtime(server_url: str, server_token: str, monkeypatch: pyte info = await client.get_cluster_info() assert info["worker_protocol"]["server_capabilities"]["prepared_local_activities"] is True assert info["worker_protocol"]["server_capabilities"]["prepared_local_activity_groups"] is True + assert info["worker_protocol"]["server_capabilities"]["prepared_local_activity_cancellation_policies"] == [ + "try_cancel", "wait_cancellation_completed", + ] assert info["worker_protocol"]["version"] == "1.20" monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20") @@ -54,13 +57,16 @@ async def prepared_local(marker: str, phase: str) -> dict[str, Any]: @workflow.defn(name="tests.python-prepared-cancellation") class PreparedCancellationWorkflow: - def run(self, ctx: Any, marker: str, group: bool = False) -> Any: + def run(self, ctx: Any, marker: str, group: bool = False, policy: str | None = None) -> Any: try: if group: - yield [ctx.local_activity("tests.python-prepared-blocked", [marker, "work-0", None]), - ctx.local_activity("tests.python-prepared-blocked", [marker, "work-1", None])] + yield [ctx.local_activity("tests.python-prepared-blocked", [marker, "work-0", None], + cancellation_policy=policy), + ctx.local_activity("tests.python-prepared-blocked", [marker, "work-1", None], + cancellation_policy=policy)] else: - yield ctx.local_activity("tests.python-prepared-blocked", [marker, "work", None]) + yield ctx.local_activity("tests.python-prepared-blocked", [marker, "work", None], + cancellation_policy=policy) except WorkflowCancelled as error: assert error.context is ctx.cancellation_context and error.context is not None delivery_path = Path(marker + ".remaining-delivery-" + str(os.getpid()) + ".json") @@ -72,11 +78,13 @@ def run(self, ctx: Any, marker: str, group: bool = False) -> Any: yield [ctx.local_activity( "tests.python-prepared-blocked", [marker, "cleanup-" + str(index), error.request_id], retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, + cancellation_policy=policy, ) for index in range(2)] else: yield ctx.local_activity( "tests.python-prepared-blocked", [marker, "cleanup", error.request_id], retry_policy={"max_attempts": 2, "backoff_seconds": [0]}, + cancellation_policy=policy, ) final_path = Path(marker + ".remaining-final.json") pending = final_path.with_suffix(".writing") @@ -229,8 +237,10 @@ async def test_prepared_nested_mixed_group_starts_peers_concurrently_and_cold_re @pytest.mark.parametrize("group", [False, True], ids=["sequential", "atomic-group"]) +@pytest.mark.parametrize("policy", [None, "try_cancel", "wait_cancellation_completed"], + ids=["historical-omission", "try-cancel", "wait-cancellation-completed"]) async def test_prepared_callback_stop_cleanup_sigkill_and_cold_recovery_keep_original_30_second_deadline( - server_url: str, server_token: str, tmp_path: Path, group: bool, + server_url: str, server_token: str, tmp_path: Path, group: bool, policy: str | None, ) -> None: queue = "py-prepared-cancel-" + uuid.uuid4().hex[:8] marker = str(tmp_path / "callback") @@ -251,7 +261,7 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: async with Client(server_url, token=server_token, namespace="default") as client: try: handle = await client.start_workflow(workflow_type="tests.python-prepared-cancellation", task_queue=queue, - workflow_id=queue, input=[marker, group]) + workflow_id=queue, input=[marker, group, policy]) first = await owner(queue + "-first", "hold") work_phases = ["work-0", "work-1"] if group else ["work"] cleanup_phases = ["cleanup-0", "cleanup-1"] if group else ["cleanup"] @@ -301,6 +311,14 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: assert kinds.count("ActivityCancellationAcknowledged") == len(work) assert kinds.count("ActivityCancelled") == len(work) assert "ActivityHeartbeatRecorded" not in kinds + scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"] + assert len(scheduled) == len(work) + len(cleanup) + assert all(event["payload"]["activity"]["cancellation_policy"] == (policy or "try_cancel") + for event in scheduled) + if policy == "wait_cancellation_completed": + delivery_index = kinds.index("CooperativeCancellationDelivered") + assert all(index < delivery_index for index, kind in enumerate(kinds) + if kind == "ActivityCancellationAcknowledged") terminal = next(event for event in history if event["event_type"] == "WorkflowCancelled") assert datetime.fromisoformat(terminal["timestamp"].replace("Z", "+00:00")) < deadline receipts = [json.loads(line) for line in trace_path.read_text().splitlines()] diff --git a/tests/test_prepared_local_activity_policies.py b/tests/test_prepared_local_activity_policies.py new file mode 100644 index 0000000..2d9c9ff --- /dev/null +++ b/tests/test_prepared_local_activity_policies.py @@ -0,0 +1,206 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import serializer, workflow +from durable_workflow.cancellation import CancellationPolicy, ParentClosePolicy +from durable_workflow.errors import NonDeterministicReplayError +from durable_workflow.worker import Worker +from durable_workflow.workflow import ( + CompleteUpdate, + LocalActivityExecutionAborted, + RecordLocalActivity, + apply_update, + query_state, + replay, + validate_update, +) +from tests.test_cancellation_context import delivery, request +from tests.test_prepared_local_activity_groups import local_event +from tests.test_prepared_local_activity_worker import PreparedServer, SequentialWorkflow +from tests.test_updates import _update_accepted_event + +POLICIES = ("try_cancel", "wait_cancellation_completed") + + +@workflow.defn(name="prepared.policy-group") +class PolicyGroupWorkflow: + def run(self, ctx: Any, policy: str = "wait_cancellation_completed") -> Any: + return (yield [ctx.local_activity("prepared.group", []), + ctx.local_activity("prepared.group", [], cancellation_policy=policy)]) + + +@workflow.defn(name="prepared.policy") +class PolicyWorkflow: + def run(self, ctx: Any, policy: Any = None, group: bool = False) -> Any: + call = ctx.local_activity("prepared.callback", [], cancellation_policy=policy) + self.result = None + self.result = yield [ctx.local_activity("prepared.callback", []), [call]] if group else call + return self.result + + @workflow.query(name="result") + def result_query(self) -> Any: + return self.result + + @workflow.update("inspect") + def inspect(self) -> Any: + return self.result + + @workflow.update_validator("inspect") + def validate_inspect(self) -> Any: + return self.result + + +def history(policy: str | None, *, complete: bool = False) -> list[dict[str, Any]]: + payload: dict[str, Any] = {"sequence": 1, "activity_type": "prepared.callback", "execution_mode": "local"} + if policy is not None: + payload["activity"] = {"cancellation_policy": policy} + events = [{"event_type": kind, "payload": dict(payload)} for kind in ("ActivityScheduled", "ActivityStarted")] + if complete: + events.append({"event_type": "ActivityCompleted", "payload": { + **payload, "result": serializer.envelope(b"completed"), + }}) + return events + + +@pytest.mark.parametrize("policy", [*POLICIES, CancellationPolicy.TRY_CANCEL, + CancellationPolicy.WAIT_CANCELLATION_COMPLETED]) +def test_explicit_policy_is_preserved_before_admission(policy: Any) -> None: + outcome = replay(PolicyWorkflow, [], [policy], prepare_local_activities=True, + local_activity_cancellation_policies=POLICIES, + local_activity_executor=lambda _: pytest.fail("callback ran before admission")) + call = outcome.prepared_local_activity + assert call is not None and not call.recover + assert call.descriptor("avro")["cancellation_policy"] == policy + assert "outcome" not in call.descriptor("avro") + + +@pytest.mark.parametrize("policy", [None, *POLICIES]) +def test_cold_replay_retains_policy_and_skips_completed_callback(policy: str | None) -> None: + options = {"prepare_local_activities": True, "local_activity_cancellation_policies": POLICIES} + pending = replay(PolicyWorkflow, history(policy), [policy], **options) + call = pending.prepared_local_activity + assert call is not None and call.recover + assert call.descriptor("avro").get("cancellation_policy") == policy + completed = replay(PolicyWorkflow, history(policy, complete=True), [policy], **options) + assert completed.prepared_local_activity is None + assert completed.commands[0].result == b"completed" # type: ignore[union-attr] + assert query_state(PolicyWorkflow, history(policy, complete=True), [policy], "result", **options) == b"completed" + + +@pytest.mark.parametrize("recorded,authored", [ + (None, "wait_cancellation_completed"), ("wait_cancellation_completed", None), + ("wait_cancellation_completed", "try_cancel"), ("try_cancel", "wait_cancellation_completed"), +]) +@pytest.mark.parametrize("complete", [False, True]) +def test_changed_policy_fails_pending_and_completed_replay(recorded: str | None, authored: str | None, + complete: bool) -> None: + with pytest.raises(NonDeterministicReplayError, match="local_activity_cancellation_policy_changed"): + replay(PolicyWorkflow, history(recorded, complete=complete), [authored], prepare_local_activities=True, + local_activity_cancellation_policies=POLICIES) + + +@pytest.mark.parametrize("delivered", [False, True]) +def test_changed_policy_is_rejected_before_request_or_committed_delivery(delivered: bool) -> None: + events = [*history("wait_cancellation_completed"), request()] + if delivered: + marker = delivery() + marker["payload"]["call_kind"] = "local_activity" + events.append(marker) + with pytest.raises(NonDeterministicReplayError, match="local_activity_cancellation_policy_changed"): + replay(PolicyWorkflow, events, ["try_cancel"], run_id="run-1", prepare_local_activities=True, + local_activity_cancellation_policies=POLICIES) + + +def test_atomic_group_recovery_preserves_policy_from_canonical_nested_snapshot() -> None: + events = [] + for sequence in (1, 2): + events.append(local_event("ActivityScheduled", sequence, activity={ + "cancellation_policy": "try_cancel" if sequence == 1 else "wait_cancellation_completed", + })) + events.append(local_event("ActivityStarted", sequence)) + options = {"prepare_local_activities": True, "prepare_local_activity_groups": True, + "local_activity_cancellation_policies": POLICIES} + outcome = replay(PolicyGroupWorkflow, events, [], **options) + group = outcome.prepared_local_activity_group + assert group is not None and group.committed and all(call.recover for call in group.calls) + assert group.calls[1].descriptor("avro")["cancellation_policy"] == "wait_cancellation_completed" + with pytest.raises(NonDeterministicReplayError, match="local_activity_cancellation_policy_changed"): + replay(PolicyGroupWorkflow, events, ["try_cancel"], **options) + + +@pytest.mark.parametrize("policy", ["abandon", CancellationPolicy.ABANDON, "unknown", 1, False, + ParentClosePolicy.ABANDON]) +def test_local_abandon_and_malformed_authoring_are_refused(policy: Any) -> None: + with pytest.raises(ValueError, match="local activity"): + RecordLocalActivity("prepared.callback", [], cancellation_policy=policy) + + +@pytest.mark.parametrize("prepared,group", [(False, False), (False, True), (True, False), (True, True)]) +def test_unsupported_policy_refuses_the_entire_call_before_callback_or_checkpoint(prepared: bool, group: bool) -> None: + with pytest.raises(LocalActivityExecutionAborted, match="installed Server discovery"): + replay(PolicyWorkflow, [], ["wait_cancellation_completed", group], prepare_local_activities=prepared, + prepare_local_activity_groups=prepared, + local_activity_executor=lambda _: pytest.fail("unsupported callback ran")) + with pytest.raises(LocalActivityExecutionAborted, match="explicit policies require prepared admission"): + RecordLocalActivity("prepared.callback", [], cancellation_policy="try_cancel").to_server_command("queue") + + +@pytest.mark.parametrize("advertised,accepted", [(None, ()), (True, ()), ("try_cancel", ()), + (["abandon", "unknown"], ()), + (["try_cancel"], ("try_cancel",)), (list(POLICIES), POLICIES)]) +async def test_worker_registers_only_discovered_prepared_policies(monkeypatch: pytest.MonkeyPatch, + advertised: Any, accepted: tuple[str, ...]) -> None: + server = PreparedServer() + capabilities = server.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"] + capabilities["prepared_local_activity_cancellation_policies"] = advertised + worker = await server.worker(monkeypatch) + assert worker._local_activity_cancellation_policies == accepted + registration = server.client.register_worker.await_args.kwargs + key = "prepared_local_activity_cancellation_policies" + assert (key in registration["capabilities"]) is bool(accepted) + assert (key in registration["capability_manifest"]) is bool(accepted) + if not accepted: + server.client.register_worker.reset_mock() + worker.capabilities = (*worker.capabilities, key) + with pytest.raises(RuntimeError, match="installed prepared policies"): + await worker._register() + server.client.register_worker.assert_not_awaited() + + +def test_explicit_worker_policy_capability_requires_prepared_admission() -> None: + with pytest.raises(ValueError, match="require prepared_local_activities"): + Worker(PreparedServer().client, task_queue="queue", + capabilities=["prepared_local_activity_cancellation_policies"]) + + +async def test_worker_refuses_undiscovered_policy_before_prefix_checkpoint_and_callback( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, +) -> None: + server = PreparedServer() + worker = await server.worker(monkeypatch) + marker = tmp_path / "callback" + task = server.task(marker) + task["arguments"] = serializer.envelope([str(marker), True, "wait_cancellation_completed"]) + assert await worker._run_workflow_task(task) is None + assert not marker.exists() and server.trace == [] + server.client.prepared_local_activity_operation.assert_not_awaited() + server.client.complete_workflow_task.assert_not_awaited() + server.client.fail_workflow_task.assert_not_awaited() + assert worker.workflows["prepared.sequential"] is SequentialWorkflow + + +@pytest.mark.parametrize("complete", [False, True]) +def test_queries_updates_and_validators_replay_explicit_local_policy_without_running_callback(complete: bool) -> None: + policy = "wait_cancellation_completed" + events = history(policy, complete=complete) + options = {"prepare_local_activities": True, "local_activity_cancellation_policies": POLICIES} + expected = b"completed" if complete else None + assert query_state(PolicyWorkflow, events, [policy], "result", **options) == expected + assert validate_update(PolicyWorkflow, events, [policy], "inspect", [], **options) == expected + updated = apply_update(PolicyWorkflow, [*events, _update_accepted_event("inspect-1", "inspect", [])], + [policy], "inspect-1", **options) + assert isinstance(updated, CompleteUpdate) and updated.result == expected diff --git a/tests/test_prepared_local_activity_worker.py b/tests/test_prepared_local_activity_worker.py index 096032f..bd3db7f 100644 --- a/tests/test_prepared_local_activity_worker.py +++ b/tests/test_prepared_local_activity_worker.py @@ -27,10 +27,11 @@ @workflow.defn(name="prepared.sequential") class SequentialWorkflow: - def run(self, ctx: workflow.WorkflowContext, marker: str, prefix: bool = False): # type: ignore[no-untyped-def] + def run(self, ctx: workflow.WorkflowContext, marker: str, prefix: bool = False, + policy: str | None = None): # type: ignore[no-untyped-def] if prefix: yield ctx.upsert_memo({"before": "local"}) - return (yield ctx.local_activity("prepared.callback", [marker])) + return (yield ctx.local_activity("prepared.callback", [marker], cancellation_policy=policy)) @workflow.defn(name="prepared.cleanup") @@ -288,6 +289,8 @@ async def operation(self, **kwargs: Any) -> dict[str, Any]: self.receipt["activity_attempt_id"] = "" payload = {"sequence": body["sequence"], "activity_type": "prepared.callback", "execution_mode": "local", "activity_execution_id": "execution", "activity_attempt_id": "attempt"} + if "cancellation_policy" in body["descriptor"]: + payload["activity"] = {"cancellation_policy": body["descriptor"]["cancellation_policy"]} self.history.extend([{"id": "scheduled", "event_type": "ActivityScheduled", "payload": payload}, {"id": "started", "event_type": "ActivityStarted", "payload": payload}]) return self.receipt @@ -339,15 +342,22 @@ def task(self, marker: Path) -> dict[str, Any]: "arguments": serializer.envelope([str(marker)]), "history_events": deepcopy(self.history)} +@pytest.mark.parametrize("policy", [None, "try_cancel", "wait_cancellation_completed"]) async def test_worker_runs_only_an_admitted_process_then_replays_canonical_result( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, policy: str | None, ) -> None: server = PreparedServer() + server.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"].update( + prepared_local_activity_cancellation_policies=["try_cancel", "wait_cancellation_completed"], + ) worker = await server.worker(monkeypatch) - commands = await worker._run_workflow_task(server.task(tmp_path / "callback")) + task = server.task(tmp_path / "callback") + task["arguments"] = serializer.envelope([str(tmp_path / "callback"), False, policy]) + commands = await worker._run_workflow_task(task) assert commands is not None and [command["type"] for command in commands] == ["complete_workflow"] assert serializer.decode_envelope(commands[0]["result"]) == {"value": b"\x00\xff", "attempt": "attempt"} assert server.trace == ["prepare", "control", "control", "outcome", "history"] + assert server.history[0]["payload"].get("activity", {}).get("cancellation_policy") == policy assert not worker._prepared_local_activity_processes server.client.fail_workflow_task.assert_not_awaited() await wait_for_exit(int((tmp_path / "callback").read_text())) @@ -395,14 +405,19 @@ async def test_lost_or_noncanonical_outcome_abandons_without_reexecuting_callbac assert not worker._prepared_local_activity_processes +@pytest.mark.parametrize("policy", [None, "try_cancel", "wait_cancellation_completed"]) async def test_control_stops_and_joins_a_gil_blocked_local_callback_without_application_heartbeats( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, policy: str | None, ) -> None: server = PreparedServer() + server.client.get_cluster_info.return_value["worker_protocol"]["server_capabilities"].update( + prepared_local_activity_cancellation_policies=["try_cancel", "wait_cancellation_completed"], + ) worker = await server.worker(monkeypatch, handler=blocked_without_python_progress) marker = tmp_path / "callback" task = server.task(marker) - outcome = replay(SequentialWorkflow, [], [str(marker)], prepare_local_activities=True) + outcome = replay(SequentialWorkflow, [], [str(marker), False, policy], prepare_local_activities=True, + local_activity_cancellation_policies=worker._local_activity_cancellation_policies) pending = asyncio.create_task(worker._execute_prepared_local_activity(task, [], outcome)) try: await wait_for_file(marker) From fce32d549ad21e0802d9a1364e1972ffb2b17dcd Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sat, 3 Oct 2026 01:18:26 +0000 Subject: [PATCH 35/41] Accept historical local policy omission in qualification --- tests/integration/test_prepared_local_activity.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index b5d78a9..a147372 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -313,7 +313,7 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: assert "ActivityHeartbeatRecorded" not in kinds scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"] assert len(scheduled) == len(work) + len(cleanup) - assert all(event["payload"]["activity"]["cancellation_policy"] == (policy or "try_cancel") + assert all(event["payload"]["activity"].get("cancellation_policy", "try_cancel") == (policy or "try_cancel") for event in scheduled) if policy == "wait_cancellation_completed": delivery_index = kinds.index("CooperativeCancellationDelivered") From d41e76deb76d3a04c25c5879984b85bb9c23cda3 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sat, 3 Oct 2026 01:18:45 +0000 Subject: [PATCH 36/41] Format historical policy qualification assertion --- tests/integration/test_prepared_local_activity.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/integration/test_prepared_local_activity.py b/tests/integration/test_prepared_local_activity.py index a147372..9c0e391 100644 --- a/tests/integration/test_prepared_local_activity.py +++ b/tests/integration/test_prepared_local_activity.py @@ -313,8 +313,10 @@ async def owner(name: str, mode: str) -> asyncio.subprocess.Process: assert "ActivityHeartbeatRecorded" not in kinds scheduled = [event for event in history if event["event_type"] == "ActivityScheduled"] assert len(scheduled) == len(work) + len(cleanup) - assert all(event["payload"]["activity"].get("cancellation_policy", "try_cancel") == (policy or "try_cancel") - for event in scheduled) + assert all( + event["payload"]["activity"].get("cancellation_policy", "try_cancel") == (policy or "try_cancel") + for event in scheduled + ) if policy == "wait_cancellation_completed": delivery_index = kinds.index("CooperativeCancellationDelivered") assert all(index < delivery_index for index, kind in enumerate(kinds) From 928ded14c45a34fa05fbb6ff1765f4820453aca9 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sat, 3 Oct 2026 08:29:27 +0000 Subject: [PATCH 37/41] Keep configured poll windows separate from HTTP grace --- src/durable_workflow/worker.py | 38 +++---------------- .../test_cooperative_cancellation.py | 2 +- tests/test_worker.py | 15 ++++++-- 3 files changed, 17 insertions(+), 38 deletions(-) diff --git a/src/durable_workflow/worker.py b/src/durable_workflow/worker.py index 6029d9d..0e4da0a 100644 --- a/src/durable_workflow/worker.py +++ b/src/durable_workflow/worker.py @@ -982,30 +982,6 @@ def _server_supports_update_validation_tasks(info: dict[str, Any]) -> bool: ) -def _server_long_poll_timeout(info: dict[str, Any]) -> float | None: - worker_protocol = info.get("worker_protocol") - if not isinstance(worker_protocol, dict): - return None - - capabilities = worker_protocol.get("server_capabilities") - if not isinstance(capabilities, dict): - return None - - timeout = capabilities.get("long_poll_timeout") - if isinstance(timeout, bool): - return None - if isinstance(timeout, int | float): - return float(timeout) if timeout > 0 else None - if isinstance(timeout, str): - try: - parsed = float(timeout) - except ValueError: - return None - return parsed if parsed > 0 else None - - return None - - def _contract_version_matches(value: Any, expected: int) -> bool: if isinstance(value, int): return value == expected @@ -1086,7 +1062,6 @@ def __init__( raise ValueError("heartbeat_interval must be positive") self._poll_timeout = poll_timeout - self._poll_http_timeout = poll_timeout self.max_concurrent_workflow_tasks = max_concurrent_workflow_tasks self.max_concurrent_activity_tasks = max_concurrent_activity_tasks self.max_concurrent_worker_sessions = max_concurrent_worker_sessions @@ -1290,9 +1265,6 @@ async def _register(self) -> None: "multiplexed workflow/update-validation polling. Refusing registration so validated " "updates cannot be accepted without validator approval or exceed worker capacity." ) - server_long_poll_timeout = _server_long_poll_timeout(info) - if server_long_poll_timeout is not None: - self._poll_http_timeout = max(self._poll_http_timeout, server_long_poll_timeout + 5.0) log.debug( "server compatibility accepted: app_version=%s control_plane=%s worker_protocol=%s", info.get("version", "unknown"), @@ -3413,7 +3385,7 @@ async def _poll_workflow_tasks(self) -> None: task = await self.client.poll_workflow_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, task_kinds=self._workflow_poll_task_kinds(), history_page_size=WORKFLOW_HISTORY_PAGE_SIZE, @@ -3524,7 +3496,7 @@ async def _poll_activity_tasks(self) -> None: task = await self.client.poll_activity_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, ) except asyncio.CancelledError: @@ -3583,7 +3555,7 @@ async def _poll_query_tasks(self, *, client: Client | None = None, track_tasks: task = await client.poll_query_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, ) except Exception as e: @@ -4113,7 +4085,7 @@ async def _run_until_loop( task = await self.client.poll_workflow_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, task_kinds=self._workflow_poll_task_kinds(), history_page_size=WORKFLOW_HISTORY_PAGE_SIZE, @@ -4167,7 +4139,7 @@ async def _run_until_loop( task = await self.client.poll_activity_task( worker_id=self.worker_id, task_queue=self.task_queue, - timeout=self._poll_http_timeout, + timeout=self._poll_timeout, build_id=self.build_id, ) if self._stop.is_set(): diff --git a/tests/integration/test_cooperative_cancellation.py b/tests/integration/test_cooperative_cancellation.py index f92f2cf..3bf81f3 100644 --- a/tests/integration/test_cooperative_cancellation.py +++ b/tests/integration/test_cooperative_cancellation.py @@ -88,7 +88,7 @@ async def poll() -> dict[str, Any]: while True: try: task = await client.poll_workflow_task( - worker_id=worker.worker_id, task_queue=worker.task_queue, timeout=worker._poll_http_timeout, + worker_id=worker.worker_id, task_queue=worker.task_queue, timeout=worker._poll_timeout, ) except ServerError as error: delay = _poll_capacity_delay(error, "workflow_task", worker.task_queue) diff --git a/tests/test_worker.py b/tests/test_worker.py index 293b144..5176277 100644 --- a/tests/test_worker.py +++ b/tests/test_worker.py @@ -711,7 +711,14 @@ async def test_worker_without_validators_does_not_poll_validation_queue(self, mo ) @pytest.mark.asyncio - async def test_register_keeps_http_timeout_above_server_long_poll(self, mock_client: AsyncMock) -> None: + @pytest.mark.parametrize(("poll_loop", "poll_method"), [ + ("_poll_workflow_tasks", "poll_workflow_task"), + ("_poll_activity_tasks", "poll_activity_task"), + ("_poll_query_tasks", "poll_query_task"), + ]) + async def test_registration_preserves_configured_poll_window( + self, mock_client: AsyncMock, poll_loop: str, poll_method: str, + ) -> None: mock_client.get_cluster_info = AsyncMock( return_value=compatible_cluster_info( worker_protocol={ @@ -736,12 +743,12 @@ async def poll_once(**_: object) -> None: worker._stop.set() return None - mock_client.poll_workflow_task.side_effect = poll_once + getattr(mock_client, poll_method).side_effect = poll_once await worker._register() - await worker._poll_workflow_tasks() + await getattr(worker, poll_loop)() - assert mock_client.poll_workflow_task.call_args.kwargs["timeout"] == 17.0 + assert getattr(mock_client, poll_method).call_args.kwargs["timeout"] == 0.01 @pytest.mark.asyncio async def test_register_keeps_baseline_capabilities_when_server_does_not_support_query_tasks( From 3cddb865f3668163b75519ec468f5c776293c8a7 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sat, 3 Oct 2026 12:04:12 +0000 Subject: [PATCH 38/41] Refuse unqualified cancellation scope replay before Python application code --- src/durable_workflow/workflow.py | 18 ++++ ...ion-scope-unqualified-worker-rejected.json | 32 ++++++ tests/test_cancellation_scope_admission.py | 102 ++++++++++++++++++ tests/test_replay_regression_corpus.py | 17 ++- 4 files changed, 165 insertions(+), 4 deletions(-) create mode 100644 tests/fixtures/replay_regressions/cancellation-scope-unqualified-worker-rejected.json create mode 100644 tests/test_cancellation_scope_admission.py diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index ccf127c..5cba2c7 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -3828,6 +3828,23 @@ def _first_yield_failure(values: Iterable[Any]) -> ActivityFailed | ChildWorkflo return None +def _assert_cancellation_scope_replay_supported(events: list[dict[str, Any]]) -> None: + """Refuse unqualified scope execution before constructing application code.""" + for event in events: + unsupported = _history_event_type(event) in { + "CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeRequestConflicted", + } + payload = event.get("payload") + if isinstance(payload, Mapping): + for container in (payload, *(payload.get(name) for name in ("activity", "timer", "child_workflow"))): + if isinstance(container, Mapping) and "cancellation_scope_id" in container: + unsupported = unsupported or container["cancellation_scope_id"] != "root" + if unsupported: + raise LocalActivityExecutionAborted( + "cancellation_scope_execution_not_supported: Python worker cannot replay scoped cancellation history", + ) + + def _replay_state( workflow_cls: type, history_events: Iterable[dict[str, Any]], @@ -3862,6 +3879,7 @@ def _replay_state( ) from exception events = list(history_events) + _assert_cancellation_scope_replay_supported(events) cancellation = read_cancellation_history(events, run_id=run_id, observation=cancellation_request) cancellation_consumed = False cancellation_intent: CancellationDelivery | None = None diff --git a/tests/fixtures/replay_regressions/cancellation-scope-unqualified-worker-rejected.json b/tests/fixtures/replay_regressions/cancellation-scope-unqualified-worker-rejected.json new file mode 100644 index 0000000..dc917b2 --- /dev/null +++ b/tests/fixtures/replay_regressions/cancellation-scope-unqualified-worker-rejected.json @@ -0,0 +1,32 @@ +{ + "$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json", + "fixture_schema": "durable-workflow.replay-regression/v1", + "id": "cancellation-scope-unqualified-worker-rejected", + "protocol_version": "1.20", + "bindings": ["python"], + "workflow": { + "type": "golden.single-activity", + "input": ["Ada"], + "payload_codec": "avro" + }, + "history": [ + { + "event_type": "CancellationScopeOpened", + "payload": { + "schema": "durable-workflow.cancellation-scope/v1", + "workflow_run_id": "scope-admission-run", + "scope_id": "scope-one", + "parent_scope_id": "root", + "sequence": 1, + "shield_parent": false + } + } + ], + "expected_replay_error": { + "type": "LocalActivityExecutionAborted", + "message_contains": "cancellation_scope_execution_not_supported.*Python" + }, + "expected": { + "command_sequence": [] + } +} diff --git a/tests/test_cancellation_scope_admission.py b/tests/test_cancellation_scope_admission.py new file mode 100644 index 0000000..4fb4cda --- /dev/null +++ b/tests/test_cancellation_scope_admission.py @@ -0,0 +1,102 @@ +from __future__ import annotations + +from copy import deepcopy +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from durable_workflow import serializer, workflow +from durable_workflow.client import Client +from durable_workflow.errors import QueryFailed +from durable_workflow.worker import Worker +from durable_workflow.workflow import LocalActivityExecutionAborted + + +@workflow.defn(name="scope-admission-probe") +class ScopeProbe: + calls: list[str] = [] + + def __init__(self) -> None: + self.calls.append("constructed") + + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + self.calls.append("run") + yield ctx.start_timer(1) + return "done" + + @workflow.query("state") + def state(self) -> str: + self.calls.append("query") + return "state" + + +@pytest.fixture(autouse=True) +def reset_calls() -> None: + ScopeProbe.calls = [] + + +def scoped_histories() -> list[dict[str, Any]]: + histories: list[dict[str, Any]] = [ + {"event_type": name, "payload": {"sequence": 1, "scope_id": "scope-one"}} + for name in ("CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeRequestConflicted") + ] + for location in (None, "activity", "timer", "child_workflow"): + value = {"cancellation_scope_id": "scope-one"} + histories.append({"event_type": "TimerScheduled", "payload": value if location is None else {location: value}}) + for malformed in (None, True, 1, "", [], {}): + histories.append({"event_type": "TimerScheduled", "payload": {"cancellation_scope_id": malformed}}) + return histories + + +@pytest.mark.parametrize("event", scoped_histories()) +def test_unqualified_scope_replay_stops_before_constructing_or_running_workflow(event: dict[str, Any]) -> None: + with pytest.raises(LocalActivityExecutionAborted, match="cancellation_scope_execution_not_supported.*Python"): + workflow.replay(ScopeProbe, [event], [], run_id="run-one", payload_codec="avro") + assert ScopeProbe.calls == [] + + +def test_query_replay_does_not_enter_scope_code_or_its_query_handler() -> None: + with pytest.raises(QueryFailed, match="cancellation_scope_execution_not_supported.*Python"): + workflow.query_state(ScopeProbe, [scoped_histories()[0]], [], "state", run_id="run-one", payload_codec="avro") + assert ScopeProbe.calls == [] + + +@pytest.mark.parametrize("explicit_root", [False, True]) +def test_historical_omission_and_explicit_root_preserve_existing_replay(explicit_root: bool) -> None: + payload: dict[str, Any] = {"sequence": 1, "duration_seconds": 1} + if explicit_root: + payload["cancellation_scope_id"] = "root" + history = [{"event_type": name, "payload": deepcopy(payload)} for name in ("TimerScheduled", "TimerFired")] + outcome = workflow.replay(ScopeProbe, history, [], run_id="run-one", payload_codec="avro") + assert [type(command).__name__ for command in outcome.commands] == ["CompleteWorkflow"] + assert ScopeProbe.calls == ["constructed", "run"] + + +def test_scope_named_application_values_are_preserved_as_data() -> None: + @workflow.defn(name="scope-data-probe") + class DataProbe: + def run(self, ctx: workflow.WorkflowContext): # type: ignore[no-untyped-def] + return (yield ctx.side_effect(lambda: pytest.fail("Recorded application data must replay."))) + + value = {"cancellation_scope_id": "application-owned", "activity": {"cancellation_scope_id": None}} + event = {"event_type": "SideEffectRecorded", "payload": { + "sequence": 1, "result": serializer.envelope(value, codec="avro"), + }} + outcome = workflow.replay(DataProbe, [event], [], run_id="run-one", payload_codec="avro") + assert outcome.commands[0].result == value + + +async def test_worker_abandons_unsupported_scope_without_publishing_a_terminal_failure( + caplog: pytest.LogCaptureFixture, +) -> None: + client = AsyncMock(spec=Client) + worker = Worker(client, task_queue="scope", worker_id="scope-worker", workflows=[ScopeProbe]) + task = {"task_id": "scope-task", "workflow_id": "scope-instance", "run_id": "run-one", + "workflow_type": "scope-admission-probe", "workflow_task_attempt": 2, "payload_codec": "avro", + "arguments": serializer.envelope([], codec="avro"), "history_events": [scoped_histories()[0]]} + assert await worker._run_workflow_task_core(task) is None + assert ScopeProbe.calls == [] + client.complete_workflow_task.assert_not_awaited() + client.fail_workflow_task.assert_not_awaited() + assert "cancellation_scope_execution_not_supported" in caplog.text diff --git a/tests/test_replay_regression_corpus.py b/tests/test_replay_regression_corpus.py index a477770..6e39eef 100644 --- a/tests/test_replay_regression_corpus.py +++ b/tests/test_replay_regression_corpus.py @@ -11,7 +11,12 @@ from durable_workflow import Replayer, Worker, serializer, workflow from durable_workflow.client import Client, WorkflowStreamAppendItem from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled, WorkflowPayloadDecodeError -from durable_workflow.workflow import WorkflowContext, commands_to_server_commands, query_state +from durable_workflow.workflow import ( + LocalActivityExecutionAborted, + WorkflowContext, + commands_to_server_commands, + query_state, +) from tests.test_golden_history_replay import ( GoldenSagaCompensationWorkflow, GoldenSignalWaitWorkflow, @@ -416,12 +421,16 @@ def test_checked_in_replay_regression_corpus_uses_official_replayer( expected_error = expected.get("error") expected_replay_error = fixture.get("expected_replay_error") if isinstance(expected_replay_error, dict): - assert expected_replay_error.get("type") == "NonDeterministicReplayError" message = expected_replay_error.get("message_contains") - workflow_sequence = expected_replay_error.get("workflow_sequence") assert isinstance(message, str) and message - assert isinstance(workflow_sequence, int) assert expected.get("command_sequence") == [] + if expected_replay_error.get("type") == "LocalActivityExecutionAborted": + with pytest.raises(LocalActivityExecutionAborted, match=message): + _execute_fixture(fixture) + return + assert expected_replay_error.get("type") == "NonDeterministicReplayError" + workflow_sequence = expected_replay_error.get("workflow_sequence") + assert isinstance(workflow_sequence, int) with pytest.raises(NonDeterministicReplayError, match=message) as captured: _execute_fixture(fixture) assert captured.value.workflow_sequence == workflow_sequence From b086c4443dc610f00927a54b32adaf9746eefe6e Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sat, 3 Oct 2026 12:49:51 +0000 Subject: [PATCH 39/41] Refuse unqualified scoped delivery replay --- src/durable_workflow/workflow.py | 2 +- tests/test_cancellation_scope_admission.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 5cba2c7..490ec52 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -3832,7 +3832,7 @@ def _assert_cancellation_scope_replay_supported(events: list[dict[str, Any]]) -> """Refuse unqualified scope execution before constructing application code.""" for event in events: unsupported = _history_event_type(event) in { - "CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeRequestConflicted", + "CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeDelivered", "CancellationScopeRequestConflicted", } payload = event.get("payload") if isinstance(payload, Mapping): diff --git a/tests/test_cancellation_scope_admission.py b/tests/test_cancellation_scope_admission.py index 4fb4cda..4527edc 100644 --- a/tests/test_cancellation_scope_admission.py +++ b/tests/test_cancellation_scope_admission.py @@ -39,7 +39,7 @@ def reset_calls() -> None: def scoped_histories() -> list[dict[str, Any]]: histories: list[dict[str, Any]] = [ {"event_type": name, "payload": {"sequence": 1, "scope_id": "scope-one"}} - for name in ("CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeRequestConflicted") + for name in ("CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeDelivered", "CancellationScopeRequestConflicted") ] for location in (None, "activity", "timer", "child_workflow"): value = {"cancellation_scope_id": "scope-one"} From 363ade8721487a947839bb0465ee5d5d9499ecea Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sat, 3 Oct 2026 12:57:20 +0000 Subject: [PATCH 40/41] Format scoped delivery admission markers --- src/durable_workflow/workflow.py | 5 ++++- tests/test_cancellation_scope_admission.py | 7 ++++++- 2 files changed, 10 insertions(+), 2 deletions(-) diff --git a/src/durable_workflow/workflow.py b/src/durable_workflow/workflow.py index 490ec52..986a04f 100644 --- a/src/durable_workflow/workflow.py +++ b/src/durable_workflow/workflow.py @@ -3832,7 +3832,10 @@ def _assert_cancellation_scope_replay_supported(events: list[dict[str, Any]]) -> """Refuse unqualified scope execution before constructing application code.""" for event in events: unsupported = _history_event_type(event) in { - "CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeDelivered", "CancellationScopeRequestConflicted", + "CancellationScopeOpened", + "CancellationScopeRequested", + "CancellationScopeDelivered", + "CancellationScopeRequestConflicted", } payload = event.get("payload") if isinstance(payload, Mapping): diff --git a/tests/test_cancellation_scope_admission.py b/tests/test_cancellation_scope_admission.py index 4527edc..6445fe8 100644 --- a/tests/test_cancellation_scope_admission.py +++ b/tests/test_cancellation_scope_admission.py @@ -39,7 +39,12 @@ def reset_calls() -> None: def scoped_histories() -> list[dict[str, Any]]: histories: list[dict[str, Any]] = [ {"event_type": name, "payload": {"sequence": 1, "scope_id": "scope-one"}} - for name in ("CancellationScopeOpened", "CancellationScopeRequested", "CancellationScopeDelivered", "CancellationScopeRequestConflicted") + for name in ( + "CancellationScopeOpened", + "CancellationScopeRequested", + "CancellationScopeDelivered", + "CancellationScopeRequestConflicted", + ) ] for location in (None, "activity", "timer", "child_workflow"): value = {"cancellation_scope_id": "scope-one"} From 4a33482644b59bda1608406a4ce0534635338155 Mon Sep 17 00:00:00 2001 From: Durable Workflow Date: Sun, 4 Oct 2026 01:41:18 +0000 Subject: [PATCH 41/41] Preserve scoped cancellation origins in child run contexts --- docs/cooperative-cancellation-design.md | 12 ++ src/durable_workflow/__init__.py | 11 +- src/durable_workflow/cancellation.py | 193 ++++++++++++++++- .../scoped-run-cancellation-context.json | 164 ++++++++++++++ tests/test_scoped_run_cancellation_context.py | 201 ++++++++++++++++++ 5 files changed, 575 insertions(+), 6 deletions(-) create mode 100644 tests/fixtures/scoped-run-cancellation-context.json create mode 100644 tests/test_scoped_run_cancellation_context.py diff --git a/docs/cooperative-cancellation-design.md b/docs/cooperative-cancellation-design.md index 103e6d9..fe013bb 100644 --- a/docs/cooperative-cancellation-design.md +++ b/docs/cooperative-cancellation-design.md @@ -40,6 +40,18 @@ mismatched local request/run identities, cycles, invalid budgets and a delivery that changes the accepted snapshot. Older cancellation histories without rich context continue delivering cancellation with `context is None`. +Candidate context v2 exposes immutable `scope_origin`, a +`ScopedCancellationContext` with the original root context and every scope +address in order. Its `deadline` is the originating scope's budget, while +`root_deadline` retains the original global deadline. The child context's own +`deadline` may be earlier when parent authority is narrower. The immediate +parent request names the last scope hop, including multiple scopes in one run. +The parser verifies that run lineage derives from the complete tree and refuses +changed metadata, repeated addresses, reentry into an earlier run or a larger +budget. Cold replay and `remaining()` retain the original narrowed authority. +Reading this metadata does not enable scope execution. Both v1 and v2 contexts +remain readable. + ## Child policies `CancellationPolicy` and `ParentClosePolicy` are available from the package diff --git a/src/durable_workflow/__init__.py b/src/durable_workflow/__init__.py index 398e3bb..5dcb0ff 100644 --- a/src/durable_workflow/__init__.py +++ b/src/durable_workflow/__init__.py @@ -16,7 +16,14 @@ AuthCompositionContractError, parse_auth_composition_contract, ) -from .cancellation import CancellationContext, CancellationLineage, CancellationPolicy, ParentClosePolicy +from .cancellation import ( + CancellationContext, + CancellationLineage, + CancellationPolicy, + ParentClosePolicy, + ScopedCancellationContext, + ScopedCancellationLineage, +) from .client import ( BridgeAdapterOutcome, Client, @@ -240,6 +247,8 @@ "BridgeAdapterOutcome", "CancellationContext", "CancellationLineage", + "ScopedCancellationContext", + "ScopedCancellationLineage", "CancellationPolicy", "ParentClosePolicy", "ChildWorkflowRetryPolicy", diff --git a/src/durable_workflow/cancellation.py b/src/durable_workflow/cancellation.py index 7edb275..edc2c94 100644 --- a/src/durable_workflow/cancellation.py +++ b/src/durable_workflow/cancellation.py @@ -95,7 +95,8 @@ class CancellationContext: """Original request metadata and budget, visible at committed delivery. Nested metadata is immutable. A child has a distinct local request ID - while retaining its root identity, request time and cleanup deadline. + while retaining its root identity and request time. A scoped child can + narrow its cleanup deadline while preserving the original root budget. """ request_id: str @@ -110,6 +111,7 @@ class CancellationContext: cleanup_deadline_at: datetime lineage: tuple[CancellationLineage, ...] _replay_clock: Callable[[], datetime] | None = field(default=None, repr=False, compare=False) + scope_origin: ScopedCancellationContext | None = None def __post_init__(self) -> None: object.__setattr__(self, "requester", MappingProxyType(dict(self.requester))) @@ -135,8 +137,17 @@ def _with_replay_clock(self, clock: Callable[[], datetime]) -> CancellationConte @classmethod def from_dict(cls, snapshot: Mapping[str, Any]) -> CancellationContext: - if snapshot.get("schema") != "durable-workflow.cancellation-context/v1": + schema = snapshot.get("schema") + if schema not in ("durable-workflow.cancellation-context/v1", "durable-workflow.cancellation-context/v2"): raise ValueError("unsupported cancellation context schema") + scope_origin = None + if schema == "durable-workflow.cancellation-context/v2": + origin = snapshot.get("scope_origin") + if not isinstance(origin, Mapping): + raise ValueError("scoped run cancellation requires its original scope context") + scope_origin = ScopedCancellationContext.from_dict(origin) + elif "scope_origin" in snapshot or "scope_authority_deadline_at" in snapshot: + raise ValueError("legacy cancellation context cannot discard a scoped origin") requester = snapshot.get("requester") if not isinstance(requester, Mapping) or not requester: raise ValueError("cancellation requester must identify its caller") @@ -173,22 +184,27 @@ def from_dict(cls, snapshot: Mapping[str, Any]) -> CancellationContext: expected_parent = normalized[-2].request_id if len(normalized) > 1 else None if ( normalized[0] != CancellationLineage(root_request_id, root_instance_id, root_run_id) - or normalized[-1].request_id != request_id or parent_request_id != expected_parent + or normalized[-1].request_id != request_id + or scope_origin is None and parent_request_id != expected_parent ): raise ValueError("cancellation lineage does not match its request identities") requested_at = _timestamp(_text(snapshot, "requested_at")) deadline = _timestamp(_text(snapshot, "cleanup_deadline_at")) if deadline <= requested_at: raise ValueError("cancellation deadline must follow the original request") + if scope_origin is not None: + _assert_scope_origin(snapshot, scope_origin, normalized, normalized_requester, requested_at, deadline) return cls( request_id, root_request_id, root_instance_id, root_run_id, parent_request_id, reason, normalized_requester, _text(snapshot, "source"), requested_at, deadline, tuple(normalized), + scope_origin=scope_origin, ) def to_dict(self) -> dict[str, Any]: """Return detached metadata in the portable context schema.""" - return { - "schema": "durable-workflow.cancellation-context/v1", + snapshot: dict[str, Any] = { + "schema": "durable-workflow.cancellation-context/v1" if self.scope_origin is None + else "durable-workflow.cancellation-context/v2", "request_id": self.request_id, "root_request_id": self.root_request_id, "root_workflow_instance_id": self.root_workflow_instance_id, @@ -201,3 +217,170 @@ def to_dict(self) -> dict[str, Any]: "cleanup_deadline_at": self.cleanup_deadline_at.isoformat(timespec="microseconds").replace("+00:00", "Z"), "lineage": [entry.to_dict() for entry in self.lineage], } + if self.scope_origin is not None: + snapshot["scope_origin"] = self.scope_origin.to_dict() + snapshot["scope_authority_deadline_at"] = snapshot["cleanup_deadline_at"] + return snapshot + + +@dataclass(frozen=True) +class ScopedCancellationLineage: + """One immutable original scope address and its bounded cleanup budget.""" + + request_id: str + workflow_instance_id: str + workflow_run_id: str + scope_id: str + cleanup_deadline_at: datetime + + def to_dict(self) -> dict[str, str]: + return { + "request_id": self.request_id, + "workflow_instance_id": self.workflow_instance_id, + "workflow_run_id": self.workflow_run_id, + "scope_id": self.scope_id, + "cleanup_deadline_at": self.cleanup_deadline_at.isoformat(timespec="microseconds").replace("+00:00", "Z"), + } + + +@dataclass(frozen=True) +class ScopedCancellationContext: + """Original scope ancestry carried by a candidate cooperative child request. + + This is immutable metadata, not permission to enter a scope body. + """ + + root_context: CancellationContext + lineage: tuple[ScopedCancellationLineage, ...] + + def __post_init__(self) -> None: + object.__setattr__(self, "lineage", tuple(self.lineage)) + + @property + def request_id(self) -> str: + return self.lineage[-1].request_id + + @property + def parent_request_id(self) -> str | None: + return self.lineage[-2].request_id if len(self.lineage) > 1 else None + + @property + def workflow_instance_id(self) -> str: + return self.lineage[-1].workflow_instance_id + + @property + def workflow_run_id(self) -> str: + return self.lineage[-1].workflow_run_id + + @property + def scope_id(self) -> str: + return self.lineage[-1].scope_id + + @property + def root_scope_id(self) -> str: + return self.lineage[0].scope_id + + @property + def requested_at(self) -> datetime: + return self.root_context.requested_at + + @property + def root_deadline(self) -> datetime: + return self.root_context.deadline + + @property + def deadline(self) -> datetime: + return self.lineage[-1].cleanup_deadline_at + + @classmethod + def from_dict(cls, snapshot: Mapping[str, Any]) -> ScopedCancellationContext: + if set(snapshot) != {"schema", "root_context", "lineage"}: + raise ValueError("scoped cancellation context has missing or unsupported fields") + root_snapshot = snapshot["root_context"] + if (snapshot["schema"] != "durable-workflow.scoped-cancellation-context/v1" + or not isinstance(root_snapshot, Mapping) + or root_snapshot.get("schema") != "durable-workflow.cancellation-context/v1"): + raise ValueError("scoped cancellation requires its original root context") + root = CancellationContext.from_dict(root_snapshot) + if (set(root_snapshot) != set(root.to_dict()) or root.request_id != root.root_request_id + or root.parent_request_id is not None or len(root.lineage) != 1): + raise ValueError("scoped cancellation root must contain the original root request") + lineage = snapshot["lineage"] + if not isinstance(lineage, list) or not lineage: + raise ValueError("scoped cancellation lineage must contain its root address") + requests: set[str] = set() + addresses: set[tuple[str, str]] = set() + instances_by_run: dict[str, str] = {} + last_run: str | None = None + deadline = root.deadline + normalized = [] + for index, entry in enumerate(lineage): + if not isinstance(entry, Mapping) or set(entry) != { + "request_id", "workflow_instance_id", "workflow_run_id", "scope_id", "cleanup_deadline_at", + }: + raise ValueError("scoped cancellation address has missing or unsupported fields") + address = ScopedCancellationLineage( + _text(entry, "request_id"), _text(entry, "workflow_instance_id"), _text(entry, "workflow_run_id"), + _text(entry, "scope_id"), _timestamp(_text(entry, "cleanup_deadline_at")), + ) + if index == 0 and ( + address.request_id != root.request_id or address.workflow_instance_id != root.root_workflow_instance_id + or address.workflow_run_id != root.root_workflow_run_id or address.cleanup_deadline_at != root.deadline + ): + raise ValueError("scoped cancellation root address does not match its request") + key = (address.workflow_run_id, address.scope_id) + if address.request_id in requests or key in addresses: + raise ValueError("scoped cancellation cannot repeat a request or address") + if address.workflow_run_id in instances_by_run and ( + instances_by_run[address.workflow_run_id] != address.workflow_instance_id + or last_run != address.workflow_run_id + ): + raise ValueError("scoped cancellation cannot reenter or reassign an earlier run") + if address.cleanup_deadline_at <= root.requested_at or address.cleanup_deadline_at > deadline: + raise ValueError("scoped cancellation cannot extend a descendant budget") + normalized.append(address) + requests.add(address.request_id) + addresses.add(key) + instances_by_run[address.workflow_run_id] = address.workflow_instance_id + last_run = address.workflow_run_id + deadline = address.cleanup_deadline_at + return cls(root, tuple(normalized)) + + def to_dict(self) -> dict[str, Any]: + return { + "schema": "durable-workflow.scoped-cancellation-context/v1", + "root_context": self.root_context.to_dict(), + "lineage": [entry.to_dict() for entry in self.lineage], + } + + +def _assert_scope_origin( + snapshot: Mapping[str, Any], origin: ScopedCancellationContext, lineage: list[CancellationLineage], + requester: dict[str, str], requested_at: datetime, deadline: datetime, +) -> None: + root = origin.root_context + last = lineage[-1] + if ( + snapshot.get("parent_request_id") != origin.request_id + or snapshot["root_request_id"] != root.root_request_id + or snapshot["root_workflow_instance_id"] != root.root_workflow_instance_id + or snapshot["root_workflow_run_id"] != root.root_workflow_run_id + or snapshot.get("reason") != root.reason or requester != root.requester + or _text(snapshot, "source") != root.source or requested_at != root.requested_at + or deadline != _timestamp(_text(snapshot, "scope_authority_deadline_at")) or deadline > origin.deadline + or last.request_id in {entry.request_id for entry in origin.lineage} + or last.workflow_run_id in {entry.workflow_run_id for entry in origin.lineage} + ): + raise ValueError("run cancellation does not preserve its original scope context") + expected = [root.lineage[0]] + for entry in origin.lineage: + if entry.workflow_run_id == root.root_workflow_run_id: + continue + address = CancellationLineage(entry.request_id, entry.workflow_instance_id, entry.workflow_run_id) + if expected[-1].workflow_run_id == entry.workflow_run_id: + expected[-1] = address + else: + expected.append(address) + expected.append(last) + if lineage != expected: + raise ValueError("run cancellation lineage discards or replaces its original scope ancestry") diff --git a/tests/fixtures/scoped-run-cancellation-context.json b/tests/fixtures/scoped-run-cancellation-context.json new file mode 100644 index 0000000..7d9d0a4 --- /dev/null +++ b/tests/fixtures/scoped-run-cancellation-context.json @@ -0,0 +1,164 @@ +{ + "native_commit": "a51395da59ef70945f91f00581ecb6bad8765bdd", + "child": { + "schema": "durable-workflow.cancellation-context/v2", + "request_id": "child-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": "inner-request", + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:15.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + }, + { + "request_id": "child-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run" + } + ], + "scope_origin": { + "schema": "durable-workflow.scoped-cancellation-context/v1", + "root_context": { + "schema": "durable-workflow.cancellation-context/v1", + "request_id": "root-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": null, + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + } + ] + }, + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "outer", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z" + }, + { + "request_id": "inner-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "inner", + "cleanup_deadline_at": "2026-10-04T00:00:20.123456Z" + } + ] + }, + "scope_authority_deadline_at": "2026-10-04T00:00:15.123456Z" + }, + "grandchild": { + "schema": "durable-workflow.cancellation-context/v2", + "request_id": "grandchild-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": "child-scope-request", + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:12.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + }, + { + "request_id": "child-scope-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run" + }, + { + "request_id": "grandchild-request", + "workflow_instance_id": "grandchild-instance", + "workflow_run_id": "grandchild-run" + } + ], + "scope_origin": { + "schema": "durable-workflow.scoped-cancellation-context/v1", + "root_context": { + "schema": "durable-workflow.cancellation-context/v1", + "request_id": "root-request", + "root_request_id": "root-request", + "root_workflow_instance_id": "root-instance", + "root_workflow_run_id": "root-run", + "parent_request_id": null, + "reason": "maintenance", + "requester": { + "type": "operator", + "id": "operator-1" + }, + "source": "api", + "requested_at": "2026-10-04T00:00:00.123456Z", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z", + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run" + } + ] + }, + "lineage": [ + { + "request_id": "root-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "outer", + "cleanup_deadline_at": "2026-10-04T00:00:30.123456Z" + }, + { + "request_id": "inner-request", + "workflow_instance_id": "root-instance", + "workflow_run_id": "root-run", + "scope_id": "inner", + "cleanup_deadline_at": "2026-10-04T00:00:20.123456Z" + }, + { + "request_id": "child-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run", + "scope_id": "root", + "cleanup_deadline_at": "2026-10-04T00:00:15.123456Z" + }, + { + "request_id": "child-scope-request", + "workflow_instance_id": "child-instance", + "workflow_run_id": "child-run", + "scope_id": "child-scope", + "cleanup_deadline_at": "2026-10-04T00:00:12.123456Z" + } + ] + }, + "scope_authority_deadline_at": "2026-10-04T00:00:12.123456Z" + } +} diff --git a/tests/test_scoped_run_cancellation_context.py b/tests/test_scoped_run_cancellation_context.py new file mode 100644 index 0000000..76b2e82 --- /dev/null +++ b/tests/test_scoped_run_cancellation_context.py @@ -0,0 +1,201 @@ +from __future__ import annotations + +import json +from copy import deepcopy +from dataclasses import FrozenInstanceError +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import pytest + +from durable_workflow import CancellationContext, ScopedCancellationContext, serializer +from durable_workflow._cooperative_cancellation import read_cancellation_history +from durable_workflow.errors import NonDeterministicReplayError, WorkflowCancelled +from durable_workflow.workflow import CompleteWorkflow, ScheduleActivity, WorkflowContext, replay +from tests.test_cooperative_cancellation import completed_activity + + +def fixtures() -> dict[str, Any]: + return json.loads((Path(__file__).parent / "fixtures/scoped-run-cancellation-context.json").read_text()) + + +def history() -> list[dict[str, Any]]: + snapshot = fixtures()["child"] + return [ + {"event_type": "CooperativeCancellationRequested", "timestamp": "2026-10-04T00:00:05Z", "payload": { + "workflow_run_id": "child-run", "workflow_instance_id": "child-instance", + "workflow_command_id": "child-request", "reason": "maintenance", + "cleanup_deadline_at": snapshot["cleanup_deadline_at"], "cancellation": deepcopy(snapshot), + }}, + {"event_type": "CooperativeCancellationDelivered", "timestamp": "2026-10-04T00:00:08Z", "payload": { + "workflow_run_id": "child-run", "workflow_command_id": "child-request", + "sequence": 1, "call_kind": "timer", "cancellation": deepcopy(snapshot), + }}, + ] + + +@pytest.mark.parametrize("name", ["child", "grandchild"]) +def test_native_context_preserves_the_complete_tree_and_original_metadata(name: str) -> None: + snapshot = fixtures()[name] + context = CancellationContext.from_dict(snapshot) + assert context.to_dict() == snapshot + assert CancellationContext.from_dict(context.to_dict()) == context + assert context.root_request_id == "root-request" + assert context.reason == "maintenance" + assert context.source == "api" + assert context.requester == {"type": "operator", "id": "operator-1"} + assert context.requested_at == datetime(2026, 10, 4, 0, 0, 0, 123456, timezone.utc) + assert context.scope_origin is not None + assert context.scope_origin.root_deadline == datetime(2026, 10, 4, 0, 0, 30, 123456, timezone.utc) + if name == "child": + assert context.parent_request_id == "inner-request" + assert [entry.scope_id for entry in context.scope_origin.lineage] == ["outer", "inner"] + assert context.scope_origin.deadline == datetime(2026, 10, 4, 0, 0, 20, 123456, timezone.utc) + assert context.deadline == datetime(2026, 10, 4, 0, 0, 15, 123456, timezone.utc) + else: + assert context.parent_request_id == "child-scope-request" + assert [entry.scope_id for entry in context.scope_origin.lineage] == ["outer", "inner", "root", "child-scope"] + assert [entry.request_id for entry in context.lineage] == [ + "root-request", "child-scope-request", "grandchild-request", + ] + assert context.deadline == datetime(2026, 10, 4, 0, 0, 12, 123456, timezone.utc) + + +def test_avro_timezone_and_object_order_do_not_change_original_scope_metadata() -> None: + original = fixtures()["grandchild"] + snapshot = deepcopy(original) + snapshot["requested_at"] = "2026-10-03T20:00:00.123456-04:00" + snapshot["cleanup_deadline_at"] = snapshot["scope_authority_deadline_at"] = "2026-10-03T20:00:12.123456-04:00" + snapshot["requester"] = dict(reversed(list(snapshot["requester"].items()))) + snapshot["scope_origin"]["lineage"] = [ + dict(reversed(list(entry.items()))) for entry in snapshot["scope_origin"]["lineage"] + ] + decoded = serializer.decode(serializer.encode(snapshot), codec="avro") + assert CancellationContext.from_dict(decoded).to_dict() == original + + +def test_origin_and_nested_scope_addresses_are_immutable() -> None: + snapshot = fixtures()["grandchild"] + context = CancellationContext.from_dict(snapshot) + snapshot["scope_origin"]["lineage"][3]["scope_id"] = "changed" + origin = context.scope_origin + assert isinstance(origin, ScopedCancellationContext) + assert origin.scope_id == "child-scope" and origin.root_scope_id == "outer" + assert origin.request_id == "child-scope-request" and origin.parent_request_id == "child-request" + assert origin.workflow_instance_id == "child-instance" and origin.workflow_run_id == "child-run" + assert origin.requested_at == context.requested_at + with pytest.raises(FrozenInstanceError): + origin.lineage[3].scope_id = "changed" # type: ignore[misc] + with pytest.raises(TypeError): + origin.root_context.requester["id"] = "changed" # type: ignore[index] + detached = origin.to_dict() + detached["lineage"][3]["scope_id"] = "changed" + assert origin.to_dict() == fixtures()["grandchild"]["scope_origin"] + + +def test_cold_replay_keeps_original_origin_and_the_same_narrowed_cleanup_clock() -> None: + contexts: list[CancellationContext] = [] + observations: list[list[float]] = [] + + class Cleanup: + def run(self, ctx: WorkflowContext): # type: ignore[no-untyped-def] + assert ctx.cancellation_context is None + try: + yield ctx.start_timer(10) + except WorkflowCancelled as error: + assert error.context is not None and error.context is ctx.cancellation_context + assert error.context.to_dict() == fixtures()["child"] + contexts.append(error.context) + seen = [error.context.remaining()] + observations.append(seen) + with ctx.cancellation_shield(): + yield ctx.schedule_activity("cleanup", []) + seen.append(error.context.remaining()) + return {"remaining": seen, "cancellation": error.context.to_dict()} + return None + + events = history() + pending = replay(Cleanup, events, [], run_id="child-run") + assert len(pending.commands) == 1 and isinstance(pending.commands[0], ScheduleActivity) + assert observations == [[7.123456]] + events.extend([ + {**completed_activity(2, "cleanup", "cleaned"), "timestamp": "2026-10-04T00:00:12Z"}, + {"event_type": "SignalReceived", "timestamp": "2026-10-04T00:00:29Z", "payload": { + "signal_name": "future", "arguments": serializer.envelope([]), + }}, + ]) + for _replacement in range(2): + assert replay(Cleanup, events, [], run_id="child-run").commands == [CompleteWorkflow({ + "remaining": [7.123456, 3.123456], "cancellation": fixtures()["child"], + })] + for context in contexts: + with pytest.raises(RuntimeError, match="active workflow replay"): + context.remaining() + + +def test_canonical_delivery_cannot_replace_a_committed_scope_origin() -> None: + events = history() + events[1]["payload"]["cancellation"]["scope_origin"]["lineage"][1]["scope_id"] = "different" + with pytest.raises(NonDeterministicReplayError): + read_cancellation_history(events, run_id="child-run") + + +def invalid_contexts() -> list[Any]: + original = fixtures()["grandchild"] + rows = [] + for key, value in { + "schema": "durable-workflow.cancellation-context/v1", "root_request_id": "other", + "root_workflow_instance_id": "other", "root_workflow_run_id": "other", + "parent_request_id": "root-request", "reason": "other", "source": "other", + "requested_at": "2026-10-04T00:00:01.123456Z", "scope_origin": [], + "cleanup_deadline_at": "2026-10-04T00:00:15.123456Z", + "scope_authority_deadline_at": "2026-10-04T00:00:15.123456Z", + }.items(): + rows.append(pytest.param({**original, key: value}, id=key)) + for key in ["scope_origin", "scope_authority_deadline_at", "parent_request_id"]: + snapshot = deepcopy(original) + del snapshot[key] + rows.append(pytest.param(snapshot, id="missing "+key)) + for name in ["requester", "discarded hop", "widened global budget", "reused request", "reused run", + "discarded root", "widened scope budget", "run reassigned", "run reentry", "repeated scope request", + "repeated address", "recursive root", "unsupported scope field", "invalid schema type"]: + snapshot = deepcopy(original) + if name == "requester": + snapshot["requester"]["id"] = "other" + elif name == "discarded hop": + snapshot["lineage"][1]["request_id"] = "child-request" + elif name == "widened global budget": + snapshot["cleanup_deadline_at"] = snapshot["scope_authority_deadline_at"] = "2026-10-04T00:00:30.123456Z" + elif name == "reused request": + snapshot["request_id"] = snapshot["lineage"][2]["request_id"] = "inner-request" + elif name == "reused run": + snapshot["lineage"][2]["workflow_run_id"] = "child-run" + elif name == "discarded root": + snapshot["scope_origin"]["lineage"].pop(0) + elif name == "widened scope budget": + snapshot["scope_origin"]["lineage"][3]["cleanup_deadline_at"] = "2026-10-04T00:00:16.123456Z" + elif name == "run reassigned": + snapshot["scope_origin"]["lineage"][3]["workflow_instance_id"] = "other" + elif name == "run reentry": + snapshot["scope_origin"]["lineage"][3]["workflow_run_id"] = "root-run" + snapshot["scope_origin"]["lineage"][3]["workflow_instance_id"] = "root-instance" + elif name == "repeated scope request": + snapshot["scope_origin"]["lineage"][3]["request_id"] = "inner-request" + elif name == "repeated address": + snapshot["scope_origin"]["lineage"][3]["scope_id"] = "root" + elif name == "recursive root": + snapshot["scope_origin"]["root_context"]["schema"] = "durable-workflow.cancellation-context/v2" + snapshot["scope_origin"]["root_context"]["scope_origin"] = original["scope_origin"] + elif name == "unsupported scope field": + snapshot["scope_origin"]["lineage"][3]["authority"] = "unrecorded" + else: + snapshot["schema"] = [] + rows.append(pytest.param(snapshot, id=name)) + return rows + + +@pytest.mark.parametrize("snapshot", invalid_contexts()) +def test_invalid_origin_identity_or_budget_is_refused(snapshot: dict[str, Any]) -> None: + with pytest.raises(ValueError): + CancellationContext.from_dict(snapshot)