Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions changelog/fixed/sleepy-firefly.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Deliver cancellation received before a synchronous activity starts when its executor thread registers, while keeping the thread available for subsequent activities.
6 changes: 4 additions & 2 deletions temporalio/worker/_activity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
103 changes: 103 additions & 0 deletions tests/worker/test_activity_thread_cancel.py
Original file line number Diff line number Diff line change
@@ -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"
)
Loading