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
+15 -2
View File
@@ -491,16 +491,29 @@ class SchedulerService:
pass
# -- 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.
Polling and settlement run in worker threads and may be waiting for
an event-loop registration callback while holding this lock. A normal
blocking ``Lock.acquire`` from the event loop would form a circular
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):
await asyncio.sleep(0)
try:
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]:
"""Run startup recovery in a worker and join it across cancellation."""
+28 -17
View File
@@ -129,13 +129,15 @@ class SchedulerResumeGate:
recovery reconcile by attempt identity.
"""
service = self._service
await service._acquire_lock_async()
cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try:
service.run_store.clear_executing(run_id)
service._live_resumes.discard(run_id)
service._resume_tasks.pop(run_id, None)
finally:
service._lock.release()
if cancelled:
raise asyncio.CancelledError
async def fence(self, run_id: str) -> None:
"""Forget a shutdown-cancelled execution; keep durable marks.
@@ -148,12 +150,14 @@ class SchedulerResumeGate:
never silently resumable, never replayed.
"""
service = self._service
await service._acquire_lock_async()
cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try:
service._live_resumes.discard(run_id)
service._resume_tasks.pop(run_id, None)
finally:
service._lock.release()
if cancelled:
raise asyncio.CancelledError
async def reconcile_cancelled(self, run_id: str) -> None:
"""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.
"""
service = self._service
await service._acquire_lock_async()
cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try:
if service._started and not service._stopping and service._scheduler:
try:
@@ -192,6 +196,8 @@ class SchedulerResumeGate:
service._resume_tasks.pop(run_id, None)
finally:
service._lock.release()
if cancelled:
raise asyncio.CancelledError
async def note_resumed_result(
self,
@@ -212,21 +218,26 @@ class SchedulerResumeGate:
the note must never break a completed resume.
"""
service = self._service
await service._acquire_lock_async()
cancelled = await service._acquire_lock_async(continue_after_cancel=True)
try:
if not service._started:
return False
scheduler = service._scheduler
if scheduler is None:
return False
try:
return scheduler.record_resumed_stopped_result(
run_id,
status_value=status_value,
checkpoint_id=checkpoint_id,
now=service.clock(),
)
except SecondOwnerError, OSError, ValueError:
return False
noted = False
else:
scheduler = service._scheduler
if scheduler is None:
noted = False
else:
try:
noted = scheduler.record_resumed_stopped_result(
run_id,
status_value=status_value,
checkpoint_id=checkpoint_id,
now=service.clock(),
)
except SecondOwnerError, OSError, ValueError:
noted = False
finally:
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.ownership import SchedulerOwnership, SecondOwnerError
from wf_scheduling.poll import Scheduler
from wf_scheduling.resume_gate import SchedulerResumeGate
from wf_scheduling.store import FileScheduleStore
@@ -405,6 +406,57 @@ asyncio.run(main())
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:
intended = ts(2026, 9, 8, 12, 0)
service = _service(tmp_path, ScriptedRuntime("interrupt"))