3434from typing import Annotated , Any , Concatenate , Literal , ParamSpec , TypeVar , Union , get_args , get_origin , get_type_hints
3535
3636from . import serializer
37+ from ._cooperative_cancellation import CancellationRequest , read_cancellation_history
3738from .activity import ActivityContext , ActivityInfo , _set_context
3839from .auth_composition import (
3940 AUTH_COMPOSITION_CONTRACT_SCHEMA ,
4950 PROTOCOL_VERSION ,
5051 Client ,
5152 WorkflowExecution ,
53+ _protocol_version_from_env ,
54+ _supports_cooperative_cancellation_protocol ,
5255)
5356from .errors import (
5457 ActivityCancelled ,
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+
162170class _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