diff --git a/CHANGELOG.md b/CHANGELOG.md index 5a8995696..51baf99b8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -28,6 +28,9 @@ to include examples, links to docs, or any other relevant information. ### Fixed +- A cancel that arrives while an activity's result or failure is still being encoded no longer + interrupts the reporting, so the completion is still sent. + ### Security ## [1.34.0] - 2026-09-30 diff --git a/temporalio/worker/_activity.py b/temporalio/worker/_activity.py index f72df9cf4..98f79b8dd 100644 --- a/temporalio/worker/_activity.py +++ b/temporalio/worker/_activity.py @@ -197,7 +197,14 @@ async def drain_poll_queue(self) -> None: async def wait_all_completed(self) -> None: running_tasks = [v.task for v in self._running_activities.values() if v.task] if running_tasks: - await asyncio.gather(*running_tasks, return_exceptions=False) + # Never let a task exception escape and stall shutdown + for result in await asyncio.gather(*running_tasks, return_exceptions=True): + if isinstance(result, BaseException) and not isinstance( + result, asyncio.CancelledError + ): + logger.warning( + "Activity task raised during worker shutdown", exc_info=result + ) def _handle_cancel_activity_task( self, @@ -223,7 +230,7 @@ def _heartbeat(self, task_token: bytes, *details: Any) -> None: # converter is async, we have to schedule it. If the activity is done, # we do not schedule any more. Technically this should be impossible to # call if the activity is done because this sync call can only be called - # inside the activity and done is set to False when the activity + # inside the activity and done is set to True when the activity # returns. logger = temporalio.activity.logger activity = self._running_activities.get(task_token) @@ -348,9 +355,15 @@ async def _handle_start_activity_task( context, StorageDriverStoreContext(target=store_target) ) try: - result = await self._execute_activity( - start, running_activity, task_token, data_converter - ) + try: + result = await self._execute_activity( + start, running_activity, task_token, data_converter + ) + finally: + # The activity code has finished, so a cancel from here on must + # not interrupt reporting its outcome; done stops cancel() from + # cancelling this task + running_activity.done = True [payload] = await data_converter.encode([result]) completion.result.completed.result.CopyFrom(payload) except BaseException as err: @@ -472,9 +485,7 @@ async def _handle_start_activity_task( # Do final completion try: - # We mark the activity as done and let the currently running - # heartbeat task finish - running_activity.done = True + # Let the currently running heartbeat task finish if running_activity.last_heartbeat_task: try: await running_activity.last_heartbeat_task diff --git a/tests/worker/test_activity.py b/tests/worker/test_activity.py index c2add4769..9d303f9dd 100644 --- a/tests/worker/test_activity.py +++ b/tests/worker/test_activity.py @@ -1,5 +1,6 @@ import asyncio import concurrent.futures +import dataclasses import logging import logging.handlers import os @@ -30,7 +31,11 @@ WorkflowHandle, ) from temporalio.common import RawValue, RetryPolicy -from temporalio.converter import DefaultPayloadConverter +from temporalio.converter import ( + DataConverter, + DefaultPayloadConverter, + PayloadCodec, +) from temporalio.exceptions import ( ActivityError, ApplicationError, @@ -1131,6 +1136,87 @@ async def wait_on_event() -> str: assert "Worker graceful shutdown" == await handle.result() +class _BlockingEncodeCodec(PayloadCodec): + """Blocks in encode until released, like a codec waiting on a remote KMS.""" + + def __init__(self) -> None: + self.encoding_started = asyncio.Event() + self.release = asyncio.Event() + + async def encode( + self, payloads: Sequence[temporalio.api.common.v1.Payload] + ) -> list[temporalio.api.common.v1.Payload]: + self.encoding_started.set() + await self.release.wait() + return list(payloads) + + async def decode( + self, payloads: Sequence[temporalio.api.common.v1.Payload] + ) -> list[temporalio.api.common.v1.Payload]: + return list(payloads) + + +async def test_activity_cancel_during_result_encode_still_completes( + client: Client, worker: ExternalWorker, caplog: pytest.LogCaptureFixture +): + @activity.defn + async def fail_with_details() -> NoReturn: + # Details send the failure through the codec + raise ApplicationError("boom", {"detail": "x"}) + + codec = _BlockingEncodeCodec() + config = client.config() + config["data_converter"] = dataclasses.replace( + DataConverter.default, payload_codec=codec + ) + act_task_queue = str(uuid.uuid4()) + act_worker = Worker( + Client(**config), task_queue=act_task_queue, activities=[fail_with_details] + ) + run_task = asyncio.create_task(act_worker.run()) + handle = await client.start_workflow( + "kitchen_sink", + KSWorkflowParams( + actions=[ + KSAction( + execute_activity=KSExecuteActivityAction( + name="fail_with_details", + task_queue=act_task_queue, + retry_max_attempts=1, + ) + ) + ] + ), + id=str(uuid.uuid4()), + task_queue=worker.task_queue, + ) + # Shut down while the failure is still being encoded. With no grace period + # the worker cancels the activity at once, which must not keep its outcome + # from being reported, or shutdown would wait on it forever. + await codec.encoding_started.wait() + with caplog.at_level(logging.DEBUG, logger="temporalio.worker._activity"): + shutdown_task = asyncio.create_task(act_worker.shutdown()) + + # Only let encoding finish once the cancel has reached the activity, + # so it cannot complete first and let an unfixed worker pass + async def cancel_logged() -> None: + while not any( + rec.getMessage().startswith("Cancelling activity") + for rec in caplog.records + ): + await asyncio.sleep(0.01) + + await asyncio.wait_for(cancel_logged(), 20) + codec.release.set() + await asyncio.wait_for(shutdown_task, 20) + await run_task + with pytest.raises(WorkflowFailureError) as err: + await handle.result() + assert isinstance(err.value.cause, ActivityError) + assert isinstance(err.value.cause.cause, ApplicationError) + assert err.value.cause.cause.message == "boom" + + @activity.defn def picklable_wait_on_event() -> str: activity.wait_for_worker_shutdown_sync(20)