11from __future__ import annotations
22
33import asyncio
4+ import threading
5+ from concurrent .futures import ThreadPoolExecutor
46from copy import deepcopy
57from typing import Any
68from unittest .mock import AsyncMock
1113from durable_workflow .client import Client
1214from durable_workflow .errors import ServerError , WorkflowCancelled
1315from durable_workflow .worker import Worker
16+ from durable_workflow .workflow import LocalActivityExecutionAborted
1417from tests .test_cooperative_cancellation import marker , observation , request
1518from tests .test_worker import compatible_cluster_info
1619
@@ -70,8 +73,13 @@ def lease_ack(*, observed: bool = False) -> dict[str, Any]:
7073
7174
7275class 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" ])
328507async def test_invalid_or_missing_refresh_token_never_delivers (monkeypatch : pytest .MonkeyPatch , value : Any ) -> None :
329508 server = ClaimServer ()
0 commit comments