From 78865e76f3c0f155326f91216f065e9c663acc20 Mon Sep 17 00:00:00 2001 From: Yann VR Date: Sat, 10 Oct 2026 15:00:54 +0100 Subject: [PATCH] Deliver pending cancellation when an activity thread registers Prepared with OpenAI Codex. --- changelog/fixed/sleepy-firefly.md | 1 + temporalio/worker/_activity.py | 6 +- tests/worker/test_activity_thread_cancel.py | 103 ++++++++++++++++++++ 3 files changed, 108 insertions(+), 2 deletions(-) create mode 100644 changelog/fixed/sleepy-firefly.md create mode 100644 tests/worker/test_activity_thread_cancel.py diff --git a/changelog/fixed/sleepy-firefly.md b/changelog/fixed/sleepy-firefly.md new file mode 100644 index 000000000..563905160 --- /dev/null +++ b/changelog/fixed/sleepy-firefly.md @@ -0,0 +1 @@ +Deliver cancellation received before a synchronous activity starts when its executor thread registers, while keeping the thread available for subsequent activities. diff --git a/temporalio/worker/_activity.py b/temporalio/worker/_activity.py index 770b2aaa5..8e3e3d1ae 100644 --- a/temporalio/worker/_activity.py +++ b/temporalio/worker/_activity.py @@ -762,13 +762,15 @@ def __init__(self) -> None: def set_thread_id(self, thread_id: int) -> None: with self._lock: self._thread_id = thread_id + self._raise_in_thread_if_pending_unlocked() @contextmanager def active_thread(self) -> Iterator[None]: thread_id = threading.current_thread().ident - if thread_id is not None: - self.set_thread_id(thread_id) try: + # Registration can raise, so it must be covered by thread cleanup. + if thread_id is not None: + self.set_thread_id(thread_id) yield None finally: if thread_id is not None: diff --git a/tests/worker/test_activity_thread_cancel.py b/tests/worker/test_activity_thread_cancel.py new file mode 100644 index 000000000..6bb5c2550 --- /dev/null +++ b/tests/worker/test_activity_thread_cancel.py @@ -0,0 +1,103 @@ +import concurrent.futures +import threading + +import pytest + +import temporalio.exceptions +import temporalio.worker._activity + + +@pytest.mark.parametrize("cancel_before_start", [False, True]) +def test_pending_cancel_is_delivered_on_thread_registration( + cancel_before_start: bool, +) -> None: + raiser = temporalio.worker._activity._ThreadExceptionRaiser() + if cancel_before_start: + raiser.raise_in_thread(temporalio.exceptions.CancelledError) + + executed = False + + def activity() -> None: + nonlocal executed + with raiser.active_thread(): + executed = True + + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit(activity) + if cancel_before_start: + with pytest.raises(temporalio.exceptions.CancelledError): + future.result(timeout=5) + else: + future.result(timeout=5) + + assert executed is not cancel_before_start + assert raiser._thread_id is None + assert raiser._pending_exception is None + assert ( + executor.submit(lambda: "next activity").result(timeout=5) + == "next activity" + ) + + +def test_pending_cancel_respects_thread_shield() -> None: + raiser = temporalio.worker._activity._ThreadExceptionRaiser() + raiser.raise_in_thread(temporalio.exceptions.CancelledError) + executed = False + + def activity() -> None: + nonlocal executed + with raiser.active_thread(): + executed = True + assert raiser._pending_exception is temporalio.exceptions.CancelledError + + with raiser.shielded(): + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + executor.submit(activity).result(timeout=5) + + assert executed + assert raiser._thread_id is None + assert raiser._pending_exception is None + + +def test_cancel_during_activity_preserves_executor_thread() -> None: + raiser = temporalio.worker._activity._ThreadExceptionRaiser() + started = threading.Event() + release = threading.Event() + + def activity() -> None: + with raiser.active_thread(): + started.set() + assert release.wait(timeout=5) + + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit(activity) + try: + assert started.wait(timeout=5) + raiser.raise_in_thread(temporalio.exceptions.CancelledError) + finally: + release.set() + + with pytest.raises(temporalio.exceptions.CancelledError): + future.result(timeout=5) + + assert raiser._thread_id is None + assert ( + executor.submit(lambda: "next activity").result(timeout=5) + == "next activity" + ) + + +def test_cancel_after_activity_preserves_executor_thread() -> None: + raiser = temporalio.worker._activity._ThreadExceptionRaiser() + + def activity() -> str: + with raiser.active_thread(): + return "first activity" + + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + assert executor.submit(activity).result(timeout=5) == "first activity" + raiser.raise_in_thread(temporalio.exceptions.CancelledError) + assert ( + executor.submit(lambda: "next activity").result(timeout=5) + == "next activity" + )