Skip to content

Commit 797e393

Browse files
Drain cooperative local cleanup and fence it at shutdown deadline
1 parent febb590 commit 797e393

2 files changed

Lines changed: 220 additions & 8 deletions

File tree

‎src/durable_workflow/worker.py‎

Lines changed: 37 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717

1818
import asyncio
1919
import contextlib
20+
import contextvars
2021
import hashlib
2122
import inspect
2223
import json
@@ -28,6 +29,7 @@
2829
import types
2930
import uuid
3031
from collections.abc import Awaitable, Callable, Coroutine, Iterable, Mapping
32+
from concurrent.futures import ThreadPoolExecutor
3133
from datetime import datetime, timezone
3234
from functools import wraps
3335
from types import FunctionType
@@ -1056,6 +1058,8 @@ def __init__(
10561058
self.max_concurrent_worker_sessions = max_concurrent_worker_sessions
10571059
self._worker_sessions: dict[str, WorkerSession] = {}
10581060
self._stop = asyncio.Event()
1061+
self._local_activity_shutdown = asyncio.Event()
1062+
self._local_activity_executor: ThreadPoolExecutor | None = None
10591063
self._wf_semaphore = asyncio.Semaphore(max_concurrent_workflow_tasks)
10601064
self._act_semaphore = asyncio.Semaphore(max_concurrent_activity_tasks)
10611065
self._shutdown_timeout = shutdown_timeout
@@ -1395,6 +1399,8 @@ async def _resolve_workflow_nexus_commands(
13951399
raise RuntimeError("workflow yielded too many consecutive Nexus service calls")
13961400

13971401
async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None:
1402+
if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set():
1403+
raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim")
13981404
task_id = str(task["task_id"])
13991405
attempt = int(task.get("workflow_task_attempt", 1))
14001406
try:
@@ -1405,6 +1411,8 @@ async def _renew_local_workflow_lease(self, task: dict[str, Any]) -> None:
14051411
)
14061412
except Exception as exc:
14071413
raise LocalActivityExecutionAborted("workflow task lease renewal failed") from exc
1414+
if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set():
1415+
raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim")
14081416
if not isinstance(response, Mapping) or any((
14091417
response.get("task_id") != task_id,
14101418
response.get("lease_owner") != self.worker_id,
@@ -1548,16 +1556,22 @@ async def observe_lease() -> None:
15481556
await self._renew_local_workflow_lease(task)
15491557

15501558
invocation = asyncio.create_task(self._execute_activity_callable(
1551-
task, command.activity_type, tuple(command.arguments), handler,
1559+
task, command.activity_type, tuple(command.arguments), handler, run_sync_in_thread=True,
15521560
))
15531561
observation = asyncio.create_task(observe_lease())
1562+
shutdown = asyncio.create_task(self._local_activity_shutdown.wait())
15541563
try:
1555-
done, _ = await asyncio.wait([invocation, observation], return_when=asyncio.FIRST_COMPLETED)
1564+
done, _ = await asyncio.wait([invocation, observation, shutdown], return_when=asyncio.FIRST_COMPLETED)
1565+
if shutdown in done:
1566+
raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim")
15561567
if observation in done:
15571568
await observation # Propagate transport observation or lost lease, never a workflow cancellation.
15581569
raise LocalActivityExecutionAborted("local lease observer stopped without an acknowledgment")
15591570
return await invocation
15601571
finally:
1572+
shutdown.cancel()
1573+
with contextlib.suppress(asyncio.CancelledError):
1574+
await shutdown
15611575
observation.cancel()
15621576
with contextlib.suppress(asyncio.CancelledError, Exception):
15631577
await observation
@@ -1651,7 +1665,10 @@ def check_boundary(
16511665
now = time.monotonic()
16521666
if state["lease_aborted"]:
16531667
raise LocalActivityExecutionAborted("local activity lost its workflow task lease")
1654-
if self._stop.is_set() or (
1668+
if self._cooperative_cancellation_supported and self._local_activity_shutdown.is_set():
1669+
state["lease_aborted"] = True
1670+
raise LocalActivityExecutionAborted("worker shutdown abandoned its local workflow claim")
1671+
if (self._stop.is_set() and not self._cooperative_cancellation_supported) or (
16551672
task.get("cancel_requested") is True and task.get("cancellation_request") is None
16561673
and task.get("_delivered_cancellation_request_id") is None
16571674
):
@@ -2320,6 +2337,7 @@ async def _execute_activity_callable(
23202337
activity_type: str,
23212338
args: tuple[Any, ...],
23222339
fn: Callable[..., Any],
2340+
*, run_sync_in_thread: bool = False,
23232341
) -> Any:
23242342
context = ActivityInterceptorContext(
23252343
worker_id=self.worker_id,
@@ -2330,7 +2348,17 @@ async def _execute_activity_callable(
23302348
)
23312349

23322350
async def call_activity(ctx: ActivityInterceptorContext) -> Any:
2333-
result = fn(*ctx.args)
2351+
if run_sync_in_thread and not inspect.iscoroutinefunction(fn):
2352+
if self._local_activity_executor is None:
2353+
# Replay threads wait for local results, so sharing their pool can deadlock.
2354+
self._local_activity_executor = ThreadPoolExecutor(
2355+
max_workers=self.max_concurrent_workflow_tasks, thread_name_prefix="dw-local-activity",
2356+
)
2357+
result = await asyncio.get_running_loop().run_in_executor(
2358+
self._local_activity_executor, contextvars.copy_context().run, fn, *ctx.args,
2359+
)
2360+
else:
2361+
result = fn(*ctx.args)
23342362
if asyncio.iscoroutine(result):
23352363
return await result
23362364
return result
@@ -3527,6 +3555,8 @@ async def _shutdown(self) -> None:
35273555
in_flight,
35283556
timeout=self._remaining_shutdown_time(deadline),
35293557
)
3558+
if pending:
3559+
self._local_activity_shutdown.set()
35303560
for t in pending:
35313561
t.cancel()
35323562
if pending:
@@ -3541,6 +3571,9 @@ async def _shutdown(self) -> None:
35413571
)
35423572
await asyncio.gather(*in_flight, return_exceptions=True)
35433573

3574+
if self._local_activity_executor is not None:
3575+
self._local_activity_executor.shutdown(wait=False, cancel_futures=True)
3576+
35443577
for session in self._worker_sessions.values():
35453578
if not session.active:
35463579
continue

‎tests/test_cooperative_cancellation_worker.py‎

Lines changed: 183 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
from __future__ import annotations
22

33
import asyncio
4+
import threading
5+
from concurrent.futures import ThreadPoolExecutor
46
from copy import deepcopy
57
from typing import Any
68
from unittest.mock import AsyncMock
@@ -11,6 +13,7 @@
1113
from durable_workflow.client import Client
1214
from durable_workflow.errors import ServerError, WorkflowCancelled
1315
from durable_workflow.worker import Worker
16+
from durable_workflow.workflow import LocalActivityExecutionAborted
1417
from tests.test_cooperative_cancellation import marker, observation, request
1518
from tests.test_worker import compatible_cluster_info
1619

@@ -70,8 +73,13 @@ def lease_ack(*, observed: bool = False) -> dict[str, Any]:
7073

7174

7275
class ClaimServer:
73-
def __init__(self, *, history: list[dict[str, Any]] | None = None) -> None:
76+
def __init__(
77+
self, *, history: list[dict[str, Any]] | None = None,
78+
lease_owner: str = "cooperative-worker", workflow_task_attempt: int = 4,
79+
) -> None:
7480
self.client = AsyncMock(spec=Client)
81+
self.lease_owner = lease_owner
82+
self.workflow_task_attempt = workflow_task_attempt
7583
self.history = list(history if history is not None else [request()])
7684
self.trace: list[str] = []
7785
self.delivery_error: Exception | None = None
@@ -83,15 +91,17 @@ def __init__(self, *, history: list[dict[str, Any]] | None = None) -> None:
8391
},
8492
})
8593
self.client.register_worker.return_value = {"registered": True}
86-
self.client.heartbeat_workflow_task.return_value = lease_ack()
94+
self.client.heartbeat_workflow_task.return_value = {
95+
**lease_ack(), "lease_owner": lease_owner, "workflow_task_attempt": workflow_task_attempt,
96+
}
8797
self.client.workflow_task_history.side_effect = self.page
8898
self.client.deliver_workflow_cancellation.side_effect = self.deliver
8999
self.client.complete_workflow_task.side_effect = self.complete
90100

91101
async def page(self, **kwargs: Any) -> dict[str, Any]:
92102
assert kwargs == {
93103
"task_id": "task-1", "next_history_page_token": "opaque-first-page",
94-
"lease_owner": "cooperative-worker", "workflow_task_attempt": 4,
104+
"lease_owner": self.lease_owner, "workflow_task_attempt": self.workflow_task_attempt,
95105
}
96106
self.trace.append("history")
97107
return {"history_events": deepcopy(self.history), "next_history_page_token": None}
@@ -115,7 +125,7 @@ async def complete(self, **kwargs: Any) -> dict[str, Any]:
115125
async def worker(self, monkeypatch: pytest.MonkeyPatch, **kwargs: Any) -> Worker:
116126
monkeypatch.setenv("DURABLE_WORKFLOW_WORKER_PROTOCOL_VERSION", "1.20")
117127
worker = Worker(
118-
self.client, task_queue="queue", worker_id="cooperative-worker",
128+
self.client, task_queue="queue", worker_id=self.lease_owner,
119129
workflows=kwargs.pop("workflows", [CancellationWorkflow]),
120130
capabilities=["cooperative_cancellation"], **kwargs,
121131
)
@@ -324,6 +334,175 @@ async def compensate(request_id: str) -> str:
324334
assert server.client.deliver_workflow_cancellation.await_count == 1
325335

326336

337+
async def test_worker_stop_drains_shielded_local_cleanup_within_its_grace_period(
338+
monkeypatch: pytest.MonkeyPatch,
339+
) -> None:
340+
server = ClaimServer()
341+
entered = asyncio.Event()
342+
release = asyncio.Event()
343+
344+
@activity.defn(name="cleanup")
345+
async def compensate(request_id: str) -> str:
346+
entered.set()
347+
await release.wait()
348+
await activity.context().heartbeat()
349+
return request_id
350+
351+
worker = await server.worker(
352+
monkeypatch, workflows=[LocalCleanupWorkflow], activities=[compensate], shutdown_timeout=1,
353+
)
354+
task = claimed_task()
355+
task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro"))
356+
execution = worker._track(worker._run_workflow_task(task))
357+
await asyncio.wait_for(entered.wait(), timeout=1)
358+
stopping = asyncio.create_task(worker.stop())
359+
await asyncio.sleep(0)
360+
assert worker._stop.is_set()
361+
release.set()
362+
commands = await asyncio.wait_for(execution, timeout=1)
363+
await asyncio.wait_for(stopping, timeout=1)
364+
assert commands is not None
365+
assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"]
366+
server.client.deliver_workflow_cancellation.assert_awaited_once()
367+
server.client.fail_workflow_task.assert_not_awaited()
368+
server.client.deregister_worker_registration.assert_awaited_once_with("cooperative-worker")
369+
370+
371+
@pytest.mark.parametrize("handler_kind", ["async", "sync"])
372+
async def test_shutdown_timeout_abandons_cleanup_and_replacement_replays_original_delivery(
373+
monkeypatch: pytest.MonkeyPatch, handler_kind: str,
374+
) -> None:
375+
server = ClaimServer()
376+
entered = asyncio.Event()
377+
discarded = asyncio.Event()
378+
loop = asyncio.get_running_loop()
379+
release_thread = threading.Event()
380+
fenced_heartbeats: list[bool] = []
381+
382+
@activity.defn(name="cleanup")
383+
async def compensate(request_id: str) -> object:
384+
assert request_id == "request-1"
385+
entered.set()
386+
try:
387+
await asyncio.Event().wait()
388+
except asyncio.CancelledError:
389+
with pytest.raises(LocalActivityExecutionAborted):
390+
await activity.context().heartbeat()
391+
fenced_heartbeats.append(True)
392+
discarded.set()
393+
return object() # A late, unencodable result must never become an activity failure.
394+
395+
@activity.defn(name="cleanup")
396+
def synchronous_compensation(request_id: str) -> object:
397+
assert request_id == "request-1"
398+
context = activity.context()
399+
assert context.info.worker_id == "cooperative-worker"
400+
loop.call_soon_threadsafe(entered.set)
401+
try:
402+
assert release_thread.wait(timeout=2)
403+
heartbeat = asyncio.run_coroutine_threadsafe(context.heartbeat(), loop)
404+
try:
405+
heartbeat.result(timeout=1)
406+
except LocalActivityExecutionAborted:
407+
fenced_heartbeats.append(True)
408+
else:
409+
fenced_heartbeats.append(False)
410+
return object()
411+
finally:
412+
loop.call_soon_threadsafe(discarded.set)
413+
414+
worker = await server.worker(
415+
monkeypatch, workflows=[LocalCleanupWorkflow],
416+
activities=[compensate if handler_kind == "async" else synchronous_compensation], shutdown_timeout=0.01,
417+
)
418+
task = claimed_task()
419+
task.update(workflow_type="worker-cancellation-local-cleanup", arguments=serializer.envelope([], codec="avro"))
420+
execution = worker._track(worker._run_workflow_task(task))
421+
await asyncio.wait_for(entered.wait(), timeout=1)
422+
original_history = deepcopy(server.history)
423+
heartbeats_before_stop = server.client.heartbeat_workflow_task.await_count
424+
await asyncio.wait_for(worker.stop(), timeout=1)
425+
release_thread.set()
426+
await asyncio.wait_for(discarded.wait(), timeout=1)
427+
assert fenced_heartbeats == [True]
428+
assert server.client.heartbeat_workflow_task.await_count == heartbeats_before_stop
429+
assert execution.cancelled() or execution.result() is None
430+
assert server.history == original_history
431+
server.client.complete_workflow_task.assert_not_awaited()
432+
server.client.fail_workflow_task.assert_not_awaited()
433+
server.client.deregister_worker_registration.assert_awaited_once_with("cooperative-worker")
434+
435+
replacement = ClaimServer(history=original_history, lease_owner="replacement-worker", workflow_task_attempt=5)
436+
cleanup_ids: list[str] = []
437+
438+
@activity.defn(name="cleanup")
439+
async def resumed_cleanup(request_id: str) -> str:
440+
cleanup_ids.append(request_id)
441+
await activity.context().heartbeat()
442+
return request_id
443+
444+
successor = await replacement.worker(monkeypatch, workflows=[LocalCleanupWorkflow], activities=[resumed_cleanup])
445+
reclaimed = deepcopy(task)
446+
reclaimed["workflow_task_attempt"] = 5
447+
reclaimed["history_events"] = deepcopy(original_history)
448+
commands = await asyncio.wait_for(successor._run_workflow_task(reclaimed), timeout=1)
449+
assert cleanup_ids == ["request-1"]
450+
assert commands is not None
451+
assert [command["type"] for command in commands] == ["record_local_activity", "complete_workflow"]
452+
assert replacement.history == original_history
453+
replacement.client.deliver_workflow_cancellation.assert_not_awaited()
454+
replacement.client.fail_workflow_task.assert_not_awaited()
455+
assert replacement.client.complete_workflow_task.await_args.kwargs["lease_owner"] == "replacement-worker"
456+
assert replacement.client.complete_workflow_task.await_args.kwargs["workflow_task_attempt"] == 5
457+
458+
459+
async def test_active_synchronous_local_call_observes_request_without_blocking_the_worker(
460+
monkeypatch: pytest.MonkeyPatch,
461+
) -> None:
462+
server = ClaimServer(history=[])
463+
entered = asyncio.Event()
464+
discarded = asyncio.Event()
465+
release = threading.Event()
466+
loop = asyncio.get_running_loop()
467+
fenced: list[bool] = []
468+
loop.set_default_executor(ThreadPoolExecutor(max_workers=1))
469+
470+
@activity.defn(name="work")
471+
def work() -> object:
472+
context = activity.context()
473+
loop.call_soon_threadsafe(entered.set)
474+
try:
475+
assert release.wait(timeout=2)
476+
heartbeat = asyncio.run_coroutine_threadsafe(context.heartbeat(), loop)
477+
try:
478+
heartbeat.result(timeout=1)
479+
except LocalActivityExecutionAborted:
480+
fenced.append(True)
481+
else:
482+
fenced.append(False)
483+
return object()
484+
finally:
485+
loop.call_soon_threadsafe(discarded.set)
486+
487+
worker = await server.worker(monkeypatch, activities=[work], heartbeat_interval=0.01)
488+
execution = worker._track(worker._run_workflow_task(claimed_task("local_activity", observed=False)))
489+
await asyncio.wait_for(entered.wait(), timeout=1)
490+
server.history.append(request())
491+
server.client.heartbeat_workflow_task.return_value = lease_ack(observed=True)
492+
try:
493+
commands = await asyncio.wait_for(execution, timeout=1)
494+
assert commands is not None and [command["type"] for command in commands] == ["start_timer"]
495+
heartbeats_after_delivery = server.client.heartbeat_workflow_task.await_count
496+
finally:
497+
release.set()
498+
await asyncio.wait_for(discarded.wait(), timeout=1)
499+
assert fenced == [True]
500+
assert server.client.heartbeat_workflow_task.await_count == heartbeats_after_delivery
501+
assert server.client.deliver_workflow_cancellation.await_args.kwargs["call_kind"] == "local_activity"
502+
server.client.fail_workflow_task.assert_not_awaited()
503+
await worker.stop()
504+
505+
327506
@pytest.mark.parametrize("value", ["", None, {}, "not-a-page"])
328507
async def test_invalid_or_missing_refresh_token_never_delivers(monkeypatch: pytest.MonkeyPatch, value: Any) -> None:
329508
server = ClaimServer()

0 commit comments

Comments
 (0)