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

1from __future__ import annotations 

2 

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 

11 

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 

19 

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 

22 

23logger: logging.Logger = get_logger("prefect.server.utilities.worker_channel") 

24 

25_WORKER_CHANNEL_CLEANUP_DISPATCH_POLL_SECONDS = 1.0 

26 

27 

28@dataclass(frozen=True) 

29class WorkerCleanupInFlight: 

30 message_id: UUID 

31 reservation_token: str 

32 lease_expires_at: DateTime 

33 

34 

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() 

54 

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 ) 

71 

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 ) 

97 

98 if event_to_wake is not None: 

99 self._set_dispatch_event(event_to_wake, loop_to_wake) 

100 

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 

109 

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 

124 

125 current_in_flight[in_flight.reservation_token] = in_flight 

126 return True 

127 

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 

143 

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 

150 

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 

163 

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 

167 

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 

172 

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) 

186 

187 if not candidates: 

188 return 

189 

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 

207 

208 if not dispatched: 

209 return 

210 

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) 

216 

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 

231 

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") 

241 

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() 

257 

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() 

266 

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 

275 

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 

298 

299 try: 

300 event.clear() 

301 await self.dispatch_available( 

302 work_pool_id=work_pool_id, 

303 cleanup_queue=cleanup_queue, 

304 ) 

305 

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) 

337 

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 ) 

357 

358 for task in pending: 

359 task.cancel() 

360 await asyncio.gather(*pending, return_exceptions=True) 

361 

362 if queue_task in done: 

363 return queue_task.result() 

364 return None 

365 

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) 

372 

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) 

380 

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 

390 

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)) 

396 

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 

404 

405 try: 

406 index = connections.index(connection) 

407 except ValueError: 

408 return 

409 

410 connections.append(connections.pop(index)) 

411 

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)] 

416 

417 @staticmethod 

418 def _worker_key(connection: WorkerChannelConnection) -> tuple[UUID, UUID, str]: 

419 return (connection.work_pool_id, connection.consumer_id, connection.worker_name) 

420 

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) 

428 

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) 

441 

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) 

451 

452 

453WORKER_CLEANUP_CONNECTION_REGISTRY = WorkerCleanupConnectionRegistry() 

454 

455 

456__all__ = [ 

457 "WORKER_CLEANUP_CONNECTION_REGISTRY", 

458 "WorkerCleanupConnectionRegistry", 

459 "WorkerCleanupInFlight", 

460]