fix: finish cancelled resume cleanup

This commit is contained in:
lda
2026-09-09 23:26:18 +07:00 Verified
parent 6a87711f89
commit 5d364f005a
3 changed files with 95 additions and 19 deletions
+14 -1
View File
@@ -491,16 +491,29 @@ class SchedulerService:
pass pass
# -- internals ----------------------------------------------------- # -- internals -----------------------------------------------------
async def _acquire_lock_async(self) -> None: async def _acquire_lock_async(self, *, continue_after_cancel: bool = False) -> bool:
"""Acquire the service lock without blocking the event loop. """Acquire the service lock without blocking the event loop.
Polling and settlement run in worker threads and may be waiting for Polling and settlement run in worker threads and may be waiting for
an event-loop registration callback while holding this lock. A normal an event-loop registration callback while holding this lock. A normal
blocking ``Lock.acquire`` from the event loop would form a circular blocking ``Lock.acquire`` from the event loop would form a circular
wait, so retry non-blocking acquisition while yielding to callbacks. wait, so retry non-blocking acquisition while yielding to callbacks.
Cleanup callers can request cancellation to be deferred until after
they have acquired and released the lock; otherwise a cancelled wait
exits immediately. The return value reports deferred cancellation.
""" """
cancelled = False
while not self._lock.acquire(blocking=False): while not self._lock.acquire(blocking=False):
try:
await asyncio.sleep(0) await asyncio.sleep(0)
except asyncio.CancelledError:
if not continue_after_cancel:
raise
# The synchronous operation guarded by this lock owns its
# cleanup boundary. Do not let cancellation strand capacity;
# propagate it after the caller has finished the operation.
cancelled = True
return cancelled
async def _recover_joined(self) -> list[str]: async def _recover_joined(self) -> list[str]:
"""Run startup recovery in a worker and join it across cancellation.""" """Run startup recovery in a worker and join it across cancellation."""
+19 -8
View File
@@ -129,13 +129,15 @@ class SchedulerResumeGate:
recovery reconcile by attempt identity. recovery reconcile by attempt identity.
""" """
service = self._service service = self._service
await service._acquire_lock_async() cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try: try:
service.run_store.clear_executing(run_id) service.run_store.clear_executing(run_id)
service._live_resumes.discard(run_id) service._live_resumes.discard(run_id)
service._resume_tasks.pop(run_id, None) service._resume_tasks.pop(run_id, None)
finally: finally:
service._lock.release() service._lock.release()
if cancelled:
raise asyncio.CancelledError
async def fence(self, run_id: str) -> None: async def fence(self, run_id: str) -> None:
"""Forget a shutdown-cancelled execution; keep durable marks. """Forget a shutdown-cancelled execution; keep durable marks.
@@ -148,12 +150,14 @@ class SchedulerResumeGate:
never silently resumable, never replayed. never silently resumable, never replayed.
""" """
service = self._service service = self._service
await service._acquire_lock_async() cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try: try:
service._live_resumes.discard(run_id) service._live_resumes.discard(run_id)
service._resume_tasks.pop(run_id, None) service._resume_tasks.pop(run_id, None)
finally: finally:
service._lock.release() service._lock.release()
if cancelled:
raise asyncio.CancelledError
async def reconcile_cancelled(self, run_id: str) -> None: async def reconcile_cancelled(self, run_id: str) -> None:
"""Reconcile caller cancellation while the scheduler stays live. """Reconcile caller cancellation while the scheduler stays live.
@@ -166,7 +170,7 @@ class SchedulerResumeGate:
shutdown; a shutdown that wins the race keeps the durable fence. shutdown; a shutdown that wins the race keeps the durable fence.
""" """
service = self._service service = self._service
await service._acquire_lock_async() cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try: try:
if service._started and not service._stopping and service._scheduler: if service._started and not service._stopping and service._scheduler:
try: try:
@@ -192,6 +196,8 @@ class SchedulerResumeGate:
service._resume_tasks.pop(run_id, None) service._resume_tasks.pop(run_id, None)
finally: finally:
service._lock.release() service._lock.release()
if cancelled:
raise asyncio.CancelledError
async def note_resumed_result( async def note_resumed_result(
self, self,
@@ -212,21 +218,26 @@ class SchedulerResumeGate:
the note must never break a completed resume. the note must never break a completed resume.
""" """
service = self._service service = self._service
await service._acquire_lock_async() cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try: try:
if not service._started: if not service._started:
return False noted = False
else:
scheduler = service._scheduler scheduler = service._scheduler
if scheduler is None: if scheduler is None:
return False noted = False
else:
try: try:
return scheduler.record_resumed_stopped_result( noted = scheduler.record_resumed_stopped_result(
run_id, run_id,
status_value=status_value, status_value=status_value,
checkpoint_id=checkpoint_id, checkpoint_id=checkpoint_id,
now=service.clock(), now=service.clock(),
) )
except SecondOwnerError, OSError, ValueError: except SecondOwnerError, OSError, ValueError:
return False noted = False
finally: finally:
service._lock.release() service._lock.release()
if cancelled:
raise asyncio.CancelledError
return noted
+52
View File
@@ -36,6 +36,7 @@ from wf_scheduling.lifecycle import (
from wf_scheduling.models import Schedule from wf_scheduling.models import Schedule
from wf_scheduling.ownership import SchedulerOwnership, SecondOwnerError from wf_scheduling.ownership import SchedulerOwnership, SecondOwnerError
from wf_scheduling.poll import Scheduler from wf_scheduling.poll import Scheduler
from wf_scheduling.resume_gate import SchedulerResumeGate
from wf_scheduling.store import FileScheduleStore from wf_scheduling.store import FileScheduleStore
@@ -405,6 +406,57 @@ asyncio.run(main())
assert "registration-stop-complete" in output assert "registration-stop-complete" in output
async def test_cancelled_resume_release_waits_for_lock_cleanup(
tmp_path: Path,
) -> None:
"""Cancellation cannot skip scheduled-resume slot cleanup."""
intended = ts(2026, 9, 8, 12, 0)
service = _service(tmp_path, ScriptedRuntime("interrupt"))
lock_held = threading.Event()
release_lock = threading.Event()
holder: threading.Thread | None = None
release_task: asyncio.Task[Any] | None = None
try:
await service.start()
service.schedule_store.create_schedule(_sched_model("a", intended))
await service.poll_once(intended + timedelta(seconds=1))
run_id = _only_run_id(service.run_store)
await _wait_for(lambda: service.live_executions == 0)
gate = SchedulerResumeGate(service)
assert await gate.acquire(run_id, owner_task=asyncio.current_task())
def hold_service_lock() -> None:
service._lock.acquire()
lock_held.set()
release_lock.wait(5)
service._lock.release()
holder = threading.Thread(target=hold_service_lock, daemon=True)
holder.start()
assert await asyncio.to_thread(lock_held.wait, 5)
release_task = asyncio.create_task(gate.release(run_id))
await asyncio.sleep(0.05)
assert not release_task.done()
release_task.cancel()
release_lock.set()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(release_task, 5)
assert not service.run_store.is_executing(run_id)
assert run_id not in service._live_resumes
assert run_id not in service._resume_tasks
finally:
release_lock.set()
if release_task is not None and not release_task.done():
release_task.cancel()
await asyncio.gather(release_task, return_exceptions=True)
if holder is not None:
holder.join(timeout=5)
await service.stop()
async def test_scheduled_interrupt_stays_resumable(tmp_path: Path) -> None: async def test_scheduled_interrupt_stays_resumable(tmp_path: Path) -> None:
intended = ts(2026, 9, 8, 12, 0) intended = ts(2026, 9, 8, 12, 0)
service = _service(tmp_path, ScriptedRuntime("interrupt")) service = _service(tmp_path, ScriptedRuntime("interrupt"))