diff --git a/distributed/deploy/spec.py b/distributed/deploy/spec.py index 49d51490d7..82ba4e6607 100644 --- a/distributed/deploy/spec.py +++ b/distributed/deploy/spec.py @@ -355,7 +355,10 @@ async def _correct_state_internal(self) -> None: to_close = set(self.workers) - set(self.worker_spec) if to_close: - if self.scheduler.status == Status.running: + if ( + self.scheduler.status == Status.running + and self.status != Status.closing + ): await self.scheduler_comm.retire_workers(workers=list(to_close)) tasks = [ asyncio.create_task(self.workers[w].close()) diff --git a/distributed/deploy/tests/test_spec_cluster.py b/distributed/deploy/tests/test_spec_cluster.py index ab1ec27e31..7aa3d987fb 100644 --- a/distributed/deploy/tests/test_spec_cluster.py +++ b/distributed/deploy/tests/test_spec_cluster.py @@ -511,6 +511,42 @@ async def test_bad_close(): assert not record +@gen_test() +async def test_correct_state_skips_retirement_while_closing(): + retired = False + + class DummyScheduler: + status = Status.running + + class DummySchedulerComm: + async def retire_workers(self, workers): + nonlocal retired + retired = True + + class DummyWorker: + def __init__(self): + self.closed = False + + async def close(self): + self.closed = True + + cluster = object.__new__(SpecCluster) + cluster._lock = asyncio.Lock() + cluster._correct_state_waiting = None + cluster.status = Status.closing + cluster.scheduler = DummyScheduler() + cluster.scheduler_comm = DummySchedulerComm() + worker = DummyWorker() + cluster.workers = {"worker": worker} + cluster.worker_spec = {} + + await cluster._correct_state_internal() + + assert not retired + assert worker.closed + assert not cluster.workers + + @gen_test() async def test_shutdown_scheduler_disabled(): async with SpecCluster(