Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/utilities/worker_channel_cleanup.py: 18%
249 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
1from __future__ import annotations
3import asyncio
4import logging
5from collections import defaultdict
6from collections.abc import AsyncIterator
7from contextlib import asynccontextmanager
8from dataclasses import dataclass
9from typing import TYPE_CHECKING
10from uuid import UUID
12from prefect.logging import get_logger
13from prefect.server.worker_communication.cleanup_queue import (
14 CleanupQueueWakeup,
15 WorkerCleanupQueue,
16)
17from prefect.types import DateTime
18from prefect.types._datetime import now
20if TYPE_CHECKING: 20 ↛ 21line 20 didn't jump to line 21 because the condition on line 20 was never true
21 from prefect.server.utilities.worker_channel import WorkerChannelConnection
23logger: logging.Logger = get_logger("prefect.server.utilities.worker_channel")
25_WORKER_CHANNEL_CLEANUP_DISPATCH_POLL_SECONDS = 1.0
28@dataclass(frozen=True)
29class WorkerCleanupInFlight:
30 message_id: UUID
31 reservation_token: str
32 lease_expires_at: DateTime
35class WorkerCleanupConnectionRegistry:
36 def __init__(self) -> None:
37 self._connections_by_work_pool_id: defaultdict[
38 UUID, list[WorkerChannelConnection]
39 ] = defaultdict(list)
40 self._dispatch_locks: defaultdict[UUID, asyncio.Lock] = defaultdict(
41 asyncio.Lock
42 )
43 self._dispatch_events_by_work_pool_id: dict[UUID, asyncio.Event] = {}
44 self._dispatch_loops_by_work_pool_id: dict[UUID, asyncio.AbstractEventLoop] = {}
45 self._dispatch_tasks_by_work_pool_id: dict[UUID, asyncio.Task[None]] = {}
46 self._exiting_dispatch_tasks_by_work_pool_id: dict[
47 UUID, asyncio.Task[None]
48 ] = {}
49 self._cleanup_in_flight_by_worker: defaultdict[
50 tuple[UUID, UUID, str], dict[str, WorkerCleanupInFlight]
51 ] = defaultdict(dict)
52 self._cleanup_dispatching_by_worker: set[tuple[UUID, UUID, str]] = set()
53 self._lock = asyncio.Lock()
55 @asynccontextmanager
56 async def register(
57 self, connection: WorkerChannelConnection
58 ) -> AsyncIterator[None]:
59 async with self._lock:
60 self._prune_expired_cleanup_reservations_for_work_pool_locked(
61 connection.work_pool_id
62 )
63 self._connections_by_work_pool_id[connection.work_pool_id].append(
64 connection
65 )
66 assert connection._cleanup_queue is not None
67 self._ensure_cleanup_dispatcher_locked(
68 work_pool_id=connection.work_pool_id,
69 cleanup_queue=connection._cleanup_queue,
70 )
72 try:
73 yield
74 finally:
75 event_to_wake: asyncio.Event | None = None
76 loop_to_wake: asyncio.AbstractEventLoop | None = None
77 async with self._lock:
78 connections = self._connections_by_work_pool_id.get(
79 connection.work_pool_id
80 )
81 if connections is not None:
82 try:
83 connections.remove(connection)
84 except ValueError:
85 pass
86 else:
87 if not connections:
88 self._connections_by_work_pool_id.pop(
89 connection.work_pool_id, None
90 )
91 event_to_wake = self._dispatch_events_by_work_pool_id.get(
92 connection.work_pool_id, None
93 )
94 loop_to_wake = self._dispatch_loops_by_work_pool_id.get(
95 connection.work_pool_id, None
96 )
98 if event_to_wake is not None:
99 self._set_dispatch_event(event_to_wake, loop_to_wake)
101 async def has_cleanup_capacity(
102 self, connection: WorkerChannelConnection, max_cleanup_concurrency: int
103 ) -> bool:
104 async with self._lock:
105 in_flight = self._cleanup_in_flight_for_connection_locked(connection)
106 self._prune_expired_cleanup_reservations_locked(in_flight)
107 self._drop_empty_cleanup_in_flight_locked(connection, in_flight)
108 return len(in_flight) < max_cleanup_concurrency
110 async def track_cleanup_reservation(
111 self,
112 connection: WorkerChannelConnection,
113 in_flight: WorkerCleanupInFlight,
114 max_cleanup_concurrency: int,
115 ) -> bool:
116 async with self._lock:
117 current_in_flight = self._cleanup_in_flight_for_connection_locked(
118 connection
119 )
120 self._prune_expired_cleanup_reservations_locked(current_in_flight)
121 if len(current_in_flight) >= max_cleanup_concurrency:
122 self._drop_empty_cleanup_in_flight_locked(connection, current_in_flight)
123 return False
125 current_in_flight[in_flight.reservation_token] = in_flight
126 return True
128 async def update_cleanup_lease(
129 self,
130 connection: WorkerChannelConnection,
131 *,
132 reservation_token: str,
133 lease_expires_at: DateTime,
134 ) -> bool:
135 async with self._lock:
136 current_in_flight = self._cleanup_in_flight_for_connection_locked(
137 connection
138 )
139 in_flight = current_in_flight.get(reservation_token)
140 if in_flight is None:
141 self._drop_empty_cleanup_in_flight_locked(connection, current_in_flight)
142 return False
144 current_in_flight[reservation_token] = WorkerCleanupInFlight(
145 message_id=in_flight.message_id,
146 reservation_token=in_flight.reservation_token,
147 lease_expires_at=lease_expires_at,
148 )
149 return True
151 async def forget_cleanup_reservation(
152 self,
153 connection: WorkerChannelConnection,
154 *,
155 message_id: UUID,
156 reservation_token: str,
157 ) -> bool:
158 async with self._lock:
159 key = self._worker_key(connection)
160 current_in_flight = self._cleanup_in_flight_by_worker.get(key)
161 if current_in_flight is None:
162 return False
164 in_flight = current_in_flight.get(reservation_token)
165 if in_flight is None or in_flight.message_id != message_id:
166 return False
168 current_in_flight.pop(reservation_token)
169 if not current_in_flight:
170 self._cleanup_in_flight_by_worker.pop(key, None)
171 return True
173 async def dispatch_available(
174 self,
175 *,
176 work_pool_id: UUID,
177 cleanup_queue: WorkerCleanupQueue,
178 ) -> None:
179 while True:
180 async with self._dispatch_locks[work_pool_id]:
181 async with self._lock:
182 self._prune_expired_cleanup_reservations_for_work_pool_locked(
183 work_pool_id
184 )
185 candidates = await self._eligible_connections(work_pool_id)
187 if not candidates:
188 return
190 dispatched = False
191 for allow_fallback_to_any_queue in (False, True):
192 for connection in candidates:
193 if not await self._claim_cleanup_dispatch(connection):
194 continue
195 try:
196 if await connection.dispatch_one_cleanup_message(
197 cleanup_queue=cleanup_queue,
198 allow_fallback_to_any_queue=allow_fallback_to_any_queue,
199 ):
200 await self._mark_cleanup_dispatched(connection)
201 dispatched = True
202 break
203 finally:
204 await self._release_cleanup_dispatch_claim(connection)
205 if dispatched:
206 break
208 if not dispatched:
209 return
211 def wake_dispatcher(self, work_pool_id: UUID) -> None:
212 event = self._dispatch_events_by_work_pool_id.get(work_pool_id)
213 loop = self._dispatch_loops_by_work_pool_id.get(work_pool_id)
214 if event is not None:
215 self._set_dispatch_event(event, loop)
217 def _ensure_cleanup_dispatcher_locked(
218 self,
219 *,
220 work_pool_id: UUID,
221 cleanup_queue: WorkerCleanupQueue,
222 ) -> None:
223 task = self._dispatch_tasks_by_work_pool_id.get(work_pool_id)
224 if task is not None and not task.done():
225 if (
226 self._exiting_dispatch_tasks_by_work_pool_id.get(work_pool_id)
227 is not task
228 ):
229 self.wake_dispatcher(work_pool_id)
230 return
232 if task is not None:
233 self._exiting_dispatch_tasks_by_work_pool_id.pop(work_pool_id, None)
234 if task.done():
235 try:
236 task.result()
237 except asyncio.CancelledError:
238 pass
239 except Exception:
240 logger.exception("Worker channel cleanup dispatcher failed")
242 loop = asyncio.get_running_loop()
243 event = asyncio.Event()
244 self._dispatch_events_by_work_pool_id[work_pool_id] = event
245 self._dispatch_loops_by_work_pool_id[work_pool_id] = loop
246 self._dispatch_tasks_by_work_pool_id[work_pool_id] = asyncio.create_task(
247 self._dispatch_loop(
248 work_pool_id=work_pool_id,
249 cleanup_queue=cleanup_queue,
250 )
251 )
252 logger.debug(
253 "Worker channel cleanup dispatcher started: work_pool_id=%s",
254 work_pool_id,
255 )
256 event.set()
258 @staticmethod
259 def _set_dispatch_event(
260 event: asyncio.Event, loop: asyncio.AbstractEventLoop | None
261 ) -> None:
262 if loop is not None and loop.is_running():
263 loop.call_soon_threadsafe(event.set)
264 return
265 event.set()
267 async def _dispatch_loop(
268 self,
269 *,
270 work_pool_id: UUID,
271 cleanup_queue: WorkerCleanupQueue,
272 ) -> None:
273 try:
274 wakeup_sequence = 0
276 while True:
277 async with self._lock:
278 has_connections = bool(
279 self._connections_by_work_pool_id.get(work_pool_id)
280 )
281 event = self._dispatch_events_by_work_pool_id.get(work_pool_id)
282 if not has_connections or event is None:
283 current_task = asyncio.current_task()
284 if (
285 current_task is not None
286 and self._dispatch_tasks_by_work_pool_id.get(work_pool_id)
287 is current_task
288 ):
289 self._exiting_dispatch_tasks_by_work_pool_id[
290 work_pool_id
291 ] = current_task
292 logger.debug(
293 "Worker channel cleanup dispatcher exiting: "
294 "work_pool_id=%s reason=no_connections",
295 work_pool_id,
296 )
297 return
299 try:
300 event.clear()
301 await self.dispatch_available(
302 work_pool_id=work_pool_id,
303 cleanup_queue=cleanup_queue,
304 )
306 wakeup = await self._wait_for_dispatch_wakeup(
307 work_pool_id=work_pool_id,
308 cleanup_queue=cleanup_queue,
309 after=wakeup_sequence,
310 event=event,
311 )
312 if wakeup is not None:
313 wakeup_sequence = wakeup.sequence
314 except asyncio.CancelledError:
315 raise
316 except Exception:
317 logger.exception("Worker channel cleanup dispatch failed")
318 await asyncio.sleep(_WORKER_CHANNEL_CLEANUP_DISPATCH_POLL_SECONDS)
319 finally:
320 current_task = asyncio.current_task()
321 async with self._lock:
322 if (
323 self._dispatch_tasks_by_work_pool_id.get(work_pool_id)
324 is current_task
325 ):
326 self._dispatch_tasks_by_work_pool_id.pop(work_pool_id, None)
327 self._exiting_dispatch_tasks_by_work_pool_id.pop(work_pool_id, None)
328 if not self._connections_by_work_pool_id.get(work_pool_id):
329 self._dispatch_events_by_work_pool_id.pop(work_pool_id, None)
330 self._dispatch_loops_by_work_pool_id.pop(work_pool_id, None)
331 self._dispatch_locks.pop(work_pool_id, None)
332 elif (
333 self._exiting_dispatch_tasks_by_work_pool_id.get(work_pool_id)
334 is current_task
335 ):
336 self._exiting_dispatch_tasks_by_work_pool_id.pop(work_pool_id, None)
338 async def _wait_for_dispatch_wakeup(
339 self,
340 *,
341 work_pool_id: UUID,
342 cleanup_queue: WorkerCleanupQueue,
343 after: int,
344 event: asyncio.Event,
345 ) -> CleanupQueueWakeup | None:
346 queue_task = asyncio.create_task(
347 cleanup_queue.wait_for_wakeup(
348 work_pool_id,
349 after=after,
350 timeout=_WORKER_CHANNEL_CLEANUP_DISPATCH_POLL_SECONDS,
351 )
352 )
353 event_task = asyncio.create_task(event.wait())
354 done, pending = await asyncio.wait(
355 {queue_task, event_task}, return_when=asyncio.FIRST_COMPLETED
356 )
358 for task in pending:
359 task.cancel()
360 await asyncio.gather(*pending, return_exceptions=True)
362 if queue_task in done:
363 return queue_task.result()
364 return None
366 async def _eligible_connections(
367 self, work_pool_id: UUID
368 ) -> tuple[WorkerChannelConnection, ...]:
369 async with self._lock:
370 connections = tuple(self._connections_by_work_pool_id.get(work_pool_id, ()))
371 dispatching = set(self._cleanup_dispatching_by_worker)
373 eligible = []
374 for connection in connections:
375 if self._worker_key(connection) in dispatching:
376 continue
377 if await connection.has_cleanup_capacity():
378 eligible.append(connection)
379 return tuple(eligible)
381 async def _claim_cleanup_dispatch(
382 self, connection: WorkerChannelConnection
383 ) -> bool:
384 async with self._lock:
385 key = self._worker_key(connection)
386 if key in self._cleanup_dispatching_by_worker:
387 return False
388 self._cleanup_dispatching_by_worker.add(key)
389 return True
391 async def _release_cleanup_dispatch_claim(
392 self, connection: WorkerChannelConnection
393 ) -> None:
394 async with self._lock:
395 self._cleanup_dispatching_by_worker.discard(self._worker_key(connection))
397 async def _mark_cleanup_dispatched(
398 self, connection: WorkerChannelConnection
399 ) -> None:
400 async with self._lock:
401 connections = self._connections_by_work_pool_id.get(connection.work_pool_id)
402 if connections is None or len(connections) < 2:
403 return
405 try:
406 index = connections.index(connection)
407 except ValueError:
408 return
410 connections.append(connections.pop(index))
412 def _cleanup_in_flight_for_connection_locked(
413 self, connection: WorkerChannelConnection
414 ) -> dict[str, WorkerCleanupInFlight]:
415 return self._cleanup_in_flight_by_worker[self._worker_key(connection)]
417 @staticmethod
418 def _worker_key(connection: WorkerChannelConnection) -> tuple[UUID, UUID, str]:
419 return (connection.work_pool_id, connection.consumer_id, connection.worker_name)
421 def _drop_empty_cleanup_in_flight_locked(
422 self,
423 connection: WorkerChannelConnection,
424 in_flight: dict[str, WorkerCleanupInFlight],
425 ) -> None:
426 if not in_flight:
427 self._cleanup_in_flight_by_worker.pop(self._worker_key(connection), None)
429 @staticmethod
430 def _prune_expired_cleanup_reservations_locked(
431 in_flight: dict[str, WorkerCleanupInFlight],
432 ) -> None:
433 current_time = now("UTC")
434 expired_tokens = [
435 token
436 for token, reservation in in_flight.items()
437 if reservation.lease_expires_at <= current_time
438 ]
439 for token in expired_tokens:
440 in_flight.pop(token, None)
442 def _prune_expired_cleanup_reservations_for_work_pool_locked(
443 self, work_pool_id: UUID
444 ) -> None:
445 for key, in_flight in tuple(self._cleanup_in_flight_by_worker.items()):
446 if key[0] != work_pool_id:
447 continue
448 self._prune_expired_cleanup_reservations_locked(in_flight)
449 if not in_flight:
450 self._cleanup_in_flight_by_worker.pop(key, None)
453WORKER_CLEANUP_CONNECTION_REGISTRY = WorkerCleanupConnectionRegistry()
456__all__ = [
457 "WORKER_CLEANUP_CONNECTION_REGISTRY",
458 "WorkerCleanupConnectionRegistry",
459 "WorkerCleanupInFlight",
460]