Skip to content

Commit febb590

Browse files
Connect cooperative cancellation delivery to Worker replay and local attempt fences
1 parent a3c909c commit febb590

2 files changed

Lines changed: 569 additions & 35 deletions

File tree

‎src/durable_workflow/worker.py‎

Lines changed: 188 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@
3434
from typing import Annotated, Any, Concatenate, Literal, ParamSpec, TypeVar, Union, get_args, get_origin, get_type_hints
3535

3636
from . import serializer
37+
from ._cooperative_cancellation import CancellationRequest, read_cancellation_history
3738
from .activity import ActivityContext, ActivityInfo, _set_context
3839
from .auth_composition import (
3940
AUTH_COMPOSITION_CONTRACT_SCHEMA,
@@ -49,6 +50,8 @@
4950
PROTOCOL_VERSION,
5051
Client,
5152
WorkflowExecution,
53+
_protocol_version_from_env,
54+
_supports_cooperative_cancellation_protocol,
5255
)
5356
from .errors import (
5457
ActivityCancelled,
@@ -87,6 +90,7 @@
8790
NexusServiceCall,
8891
RecordLocalActivity,
8992
RecordSideEffect,
93+
ReplayOutcome,
9094
UpsertMemo,
9195
apply_update,
9296
commands_to_server_commands,
@@ -159,6 +163,10 @@ def __init__(self, kind: str) -> None:
159163
self.kind = kind
160164

161165

166+
class _CooperativeCancellationObserved(LocalActivityExecutionAborted):
167+
"""Return transport observation to the worker, never to authored cleanup."""
168+
169+
162170
class _InvalidLocalActivityReport(NonRetryableError):
163171
pass
164172

@@ -1021,6 +1029,7 @@ def __init__(
10211029
}
10221030
self.activities = {_activity_name(a): a for a in activities}
10231031
self.capabilities = tuple(dict.fromkeys(capability.strip() for capability in capabilities))
1032+
self._cooperative_cancellation_supported = False
10241033
if any(not capability for capability in self.capabilities):
10251034
raise ValueError("worker capabilities must be non-empty strings")
10261035
self.worker_id = worker_id or f"py-worker-{uuid.uuid4().hex[:8]}"
@@ -1177,6 +1186,20 @@ async def _register(self) -> None:
11771186
raise RuntimeError(f"Server compatibility error: unable to read /api/cluster/info: {e}") from e
11781187

11791188
_validate_server_compatibility(info)
1189+
protocol = info.get("worker_protocol")
1190+
server_capabilities = protocol.get("server_capabilities") if isinstance(protocol, Mapping) else None
1191+
self._cooperative_cancellation_supported = (
1192+
"cooperative_cancellation" in self.capabilities
1193+
and isinstance(protocol, Mapping)
1194+
and isinstance(server_capabilities, Mapping)
1195+
and server_capabilities.get("cooperative_cancellation") is True
1196+
and _supports_cooperative_cancellation_protocol(protocol.get("version"))
1197+
and _supports_cooperative_cancellation_protocol(_protocol_version_from_env(
1198+
"DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", PROTOCOL_VERSION,
1199+
))
1200+
)
1201+
if "cooperative_cancellation" in self.capabilities and not self._cooperative_cancellation_supported:
1202+
raise RuntimeError("cooperative cancellation requires explicit compatible runtime and worker protocol 1.20")
11801203
self._query_tasks_supported = _server_supports_query_tasks(info)
11811204
self._workflow_memo_updates_supported = _server_supports_workflow_memo_updates(info)
11821205
has_update_validators = any(
@@ -1389,6 +1412,120 @@ async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None:
13891412
response.get("renewed") is not True,
13901413
)):
13911414
raise LocalActivityExecutionAborted("workflow task lease renewal was not fenced and acknowledged")
1415+
observation = response.get("cancellation_request")
1416+
if observation is not None:
1417+
self._observe_workflow_cancellation(task, observation)
1418+
if task.get("_delivered_cancellation_request_id") != task["cancellation_request"]["request_id"]:
1419+
raise _CooperativeCancellationObserved("cooperative request observed on the actual task heartbeat")
1420+
1421+
def _observe_workflow_cancellation(self, task: dict[str, Any], observation: Any) -> CancellationRequest:
1422+
if not self._cooperative_cancellation_supported or not isinstance(observation, Mapping):
1423+
raise LocalActivityExecutionAborted("workflow cancellation observation is not negotiated")
1424+
try:
1425+
current = CancellationRequest.from_observation(observation)
1426+
previous = task.get("cancellation_request")
1427+
if previous is not None and not isinstance(previous, Mapping):
1428+
raise ValueError("previous observation is not an object")
1429+
original = CancellationRequest.from_observation(previous) if previous is not None else None
1430+
except Exception as error:
1431+
raise LocalActivityExecutionAborted("workflow cancellation observation is malformed") from error
1432+
if original is not None and (original.request_id, original.requested_at, original.cleanup_deadline_at) != (
1433+
current.request_id, current.requested_at, current.cleanup_deadline_at,
1434+
):
1435+
raise LocalActivityExecutionAborted("workflow cancellation observation changed its original identity")
1436+
task["cancellation_request"] = dict(observation)
1437+
return current
1438+
1439+
async def _load_workflow_claim_history(
1440+
self, task: dict[str, Any], *, first_page_token: str | None = None,
1441+
) -> list[dict[str, Any]]:
1442+
history = [] if first_page_token is not None else list(task.get("history_events", []))
1443+
token = first_page_token if first_page_token is not None else task.get("next_history_page_token")
1444+
seen: set[str] = set()
1445+
while token is not None:
1446+
if not isinstance(token, str) or not token or token in seen:
1447+
raise LocalActivityExecutionAborted("workflow history paging did not advance its opaque token")
1448+
seen.add(token)
1449+
page = await self.client.workflow_task_history(
1450+
task_id=task["task_id"], next_history_page_token=token,
1451+
lease_owner=self.worker_id, workflow_task_attempt=task.get("workflow_task_attempt", 1),
1452+
)
1453+
if not isinstance(page, Mapping) or not isinstance(page.get("history_events"), list):
1454+
raise LocalActivityExecutionAborted("workflow history page was not acknowledged")
1455+
if any(not isinstance(event, dict) for event in page["history_events"]):
1456+
raise LocalActivityExecutionAborted("workflow history page contains a malformed event")
1457+
history.extend(page["history_events"])
1458+
token = page.get("next_history_page_token")
1459+
return history
1460+
1461+
async def _refresh_cancellation_history(
1462+
self, task: dict[str, Any], observed: CancellationRequest,
1463+
) -> list[dict[str, Any]]:
1464+
try:
1465+
return await self._load_workflow_claim_history(task, first_page_token=observed.history_refresh_page_token)
1466+
except Exception as error:
1467+
raise LocalActivityExecutionAborted(
1468+
"canonical cancellation history could not be loaded on this claim",
1469+
) from error
1470+
1471+
async def _replay_workflow_claim(
1472+
self, cls: type, task: dict[str, Any], history: list[dict[str, Any]], start_input: list[Any],
1473+
*, payload_codec: str | None, execute_local: Callable[[RecordLocalActivity], Any],
1474+
) -> tuple[ReplayOutcome, list[dict[str, Any]]]:
1475+
observation = task.get("cancellation_request")
1476+
if observation is not None:
1477+
observed = self._observe_workflow_cancellation(task, observation)
1478+
history = await self._refresh_cancellation_history(task, observed)
1479+
for _ in range(3):
1480+
state = read_cancellation_history(
1481+
history, run_id=task.get("run_id", ""), observation=task.get("cancellation_request"),
1482+
)
1483+
if state.request is not None and not self._cooperative_cancellation_supported:
1484+
raise LocalActivityExecutionAborted("canonical cancellation requires a negotiated capable worker")
1485+
if state.delivery is not None:
1486+
task["_delivered_cancellation_request_id"] = state.delivery.request_id
1487+
try:
1488+
outcome = await asyncio.to_thread(
1489+
replay, cls, history, start_input, workflow_id=task.get("workflow_id"),
1490+
run_id=task.get("run_id", ""),
1491+
workflow_command_id=(
1492+
_string_or_none(task.get("workflow_command_id")) or _string_or_none(task.get("task_id"))
1493+
),
1494+
payload_codec=payload_codec, external_storage=self.external_storage,
1495+
external_storage_cache=self.external_storage_cache,
1496+
cancel_requested=bool(task.get("cancel_requested", False)) and state.request is None,
1497+
cancellation_request=task.get("cancellation_request"), local_activity_executor=execute_local,
1498+
)
1499+
except _CooperativeCancellationObserved:
1500+
observed = self._observe_workflow_cancellation(task, task.get("cancellation_request"))
1501+
history = await self._refresh_cancellation_history(task, observed)
1502+
continue
1503+
intent = outcome.cancellation_delivery
1504+
if intent is None or outcome.commands:
1505+
# Earlier authored commands must commit on this claim first.
1506+
# Completion releases the claim; a successor replays their durable results.
1507+
return outcome, history
1508+
observed = self._observe_workflow_cancellation(task, task.get("cancellation_request"))
1509+
delivery_error: Exception | None = None
1510+
try:
1511+
await self.client.deliver_workflow_cancellation(
1512+
task_id=task["task_id"], lease_owner=self.worker_id,
1513+
workflow_task_attempt=task.get("workflow_task_attempt", 1),
1514+
request_id=intent.request_id, sequence=intent.sequence, call_kind=intent.call_kind,
1515+
sequence_span=intent.sequence_span, operation_sequence=intent.operation_sequence,
1516+
operation_sequence_span=intent.operation_sequence_span,
1517+
)
1518+
except Exception as error:
1519+
delivery_error = error
1520+
history = await self._refresh_cancellation_history(task, observed)
1521+
committed = read_cancellation_history(
1522+
history, run_id=task.get("run_id", ""), observation=task.get("cancellation_request"),
1523+
)
1524+
if committed.delivery != intent:
1525+
raise LocalActivityExecutionAborted(
1526+
"delivery was not proved by matching canonical history",
1527+
) from delivery_error
1528+
raise LocalActivityExecutionAborted("workflow cancellation replay did not converge on its canonical delivery")
13921529

13931530
def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str, Any]:
13941531
storage = self.external_storage
@@ -1401,6 +1538,42 @@ def _maybe_externalize_local_payload(self, envelope: dict[str, str]) -> dict[str
14011538
reference = store_external_payload(storage, data, codec=envelope["codec"])
14021539
return {"codec": envelope["codec"], "external_storage": reference.to_dict()}
14031540

1541+
async def _execute_cooperative_local_callable(
1542+
self, task: dict[str, Any], command: RecordLocalActivity, handler: Callable[..., Any],
1543+
attempt_state: dict[str, Any],
1544+
) -> Any:
1545+
async def observe_lease() -> None:
1546+
while True:
1547+
await asyncio.sleep(min(5.0, self._heartbeat_interval))
1548+
await self._renew_local_workflow_lease(task)
1549+
1550+
invocation = asyncio.create_task(self._execute_activity_callable(
1551+
task, command.activity_type, tuple(command.arguments), handler,
1552+
))
1553+
observation = asyncio.create_task(observe_lease())
1554+
try:
1555+
done, _ = await asyncio.wait([invocation, observation], return_when=asyncio.FIRST_COMPLETED)
1556+
if observation in done:
1557+
await observation # Propagate transport observation or lost lease, never a workflow cancellation.
1558+
raise LocalActivityExecutionAborted("local lease observer stopped without an acknowledgment")
1559+
return await invocation
1560+
finally:
1561+
observation.cancel()
1562+
with contextlib.suppress(asyncio.CancelledError, Exception):
1563+
await observation
1564+
if not invocation.done():
1565+
attempt_state["lease_aborted"] = True
1566+
invocation.cancel()
1567+
1568+
def discard_late_result(future: asyncio.Task[Any]) -> None:
1569+
if not future.cancelled():
1570+
future.exception()
1571+
1572+
# Python cannot forcibly stop a synchronous thread or a callable
1573+
# that suppresses cancellation. Its attempt is fenced and its
1574+
# eventual result cannot become a durable command.
1575+
invocation.add_done_callback(discard_late_result)
1576+
14041577
async def _execute_local_activity(
14051578
self,
14061579
task: dict[str, Any],
@@ -1478,7 +1651,10 @@ def check_boundary(
14781651
now = time.monotonic()
14791652
if state["lease_aborted"]:
14801653
raise LocalActivityExecutionAborted("local activity lost its workflow task lease")
1481-
if self._stop.is_set() or task.get("cancel_requested") is True:
1654+
if self._stop.is_set() or (
1655+
task.get("cancel_requested") is True and task.get("cancellation_request") is None
1656+
and task.get("_delivered_cancellation_request_id") is None
1657+
):
14821658
raise ActivityCancelled("local activity cancelled")
14831659
if (
14841660
command.heartbeat_timeout is not None
@@ -1541,11 +1717,16 @@ async def heartbeat(
15411717
)
15421718
_set_context(ActivityContext(info=info, client=self.client, heartbeat_callback=heartbeat))
15431719
try:
1544-
result = await self._execute_activity_callable(
1545-
task, command.activity_type, tuple(command.arguments), handler,
1546-
)
1720+
if self._cooperative_cancellation_supported:
1721+
result = await self._execute_cooperative_local_callable(task, command, handler, attempt_state)
1722+
else:
1723+
result = await self._execute_activity_callable(
1724+
task, command.activity_type, tuple(command.arguments), handler,
1725+
)
15471726
finally:
15481727
_set_context(None)
1728+
if self._cooperative_cancellation_supported:
1729+
await self._renew_local_workflow_lease(task)
15491730
check_boundary()
15501731
attempts.append({
15511732
"attempt_id": attempt_id,
@@ -1636,21 +1817,7 @@ async def _run_workflow_task_core(self, task: dict[str, Any]) -> list[dict[str,
16361817
task_id: str = task["task_id"]
16371818
attempt: int = task.get("workflow_task_attempt", 1)
16381819
wf_type: str = task.get("workflow_type", "")
1639-
history = task.get("history_events", [])
1640-
1641-
# The worker requests bounded history pages when polling. Do not replay
1642-
# an incomplete history if fetching a later page fails.
1643-
next_page_token = task.get("next_history_page_token")
1644-
while next_page_token:
1645-
page_data = await self.client.workflow_task_history(
1646-
task_id=task_id,
1647-
next_history_page_token=next_page_token,
1648-
lease_owner=self.worker_id,
1649-
workflow_task_attempt=attempt,
1650-
)
1651-
if page_data and page_data.get("history_events"):
1652-
history.extend(page_data["history_events"])
1653-
next_page_token = page_data.get("next_history_page_token") if page_data else None
1820+
history = await self._load_workflow_claim_history(task)
16541821

16551822
start_input: list[Any] = []
16561823
codec = task.get("payload_codec")
@@ -1773,22 +1940,8 @@ def execute_local(command: RecordLocalActivity) -> Any:
17731940
return future.result()
17741941

17751942
try:
1776-
outcome = await asyncio.to_thread(
1777-
replay,
1778-
cls,
1779-
history,
1780-
start_input,
1781-
workflow_id=task.get("workflow_id"),
1782-
run_id=run_id,
1783-
workflow_command_id=(
1784-
_string_or_none(task.get("workflow_command_id"))
1785-
or _string_or_none(task.get("task_id"))
1786-
),
1787-
payload_codec=codec,
1788-
external_storage=self.external_storage,
1789-
external_storage_cache=self.external_storage_cache,
1790-
cancel_requested=bool(task.get("cancel_requested", False)),
1791-
local_activity_executor=execute_local,
1943+
outcome, history = await self._replay_workflow_claim(
1944+
cls, task, history, start_input, payload_codec=codec, execute_local=execute_local,
17921945
)
17931946
except LocalActivityExecutionAborted as e:
17941947
log.warning("abandoning workflow task %s before local activity commit: %s", task_id, e)

0 commit comments

Comments
 (0)