Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/utilities/worker_channel.py: 16%

360 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.abc import Mapping 

6from datetime import datetime 

7from typing import TYPE_CHECKING, Any 

8from uuid import UUID 

9 

10from pydantic import ValidationError 

11from sqlalchemy.ext.asyncio import AsyncSession 

12from starlette.websockets import WebSocket 

13 

14import prefect.server.models as models 

15import prefect.server.schemas as schemas 

16from prefect._internal.schemas.bases import PrefectBaseModel 

17from prefect._internal.uuid7 import uuid7 

18from prefect.client.schemas.worker_channel import ( 

19 WORKER_CHANNEL_CLOSE_POLICIES, 

20 CleanupAckFrame, 

21 CleanupKind, 

22 CleanupMessageFrame, 

23 CleanupOperationResultFrame, 

24 CleanupReleaseFrame, 

25 CleanupRenewFrame, 

26 WorkerChannelCloseReason, 

27 WorkerHeartbeatFrame, 

28 WorkerReadyFrame, 

29 WorkPoolSnapshot, 

30 WorkPoolSnapshotFrame, 

31 WorkPoolSnapshotPayload, 

32 validate_worker_channel_frame, 

33) 

34from prefect.logging import get_logger 

35from prefect.server.database import PrefectDBInterface 

36from prefect.server.models.workers import emit_work_pool_status_event 

37from prefect.server.utilities import messaging, subscriptions 

38from prefect.server.utilities.worker_channel_cleanup import ( 

39 WORKER_CLEANUP_CONNECTION_REGISTRY, 

40 WorkerCleanupConnectionRegistry, 

41 WorkerCleanupInFlight, 

42) 

43from prefect.server.worker_communication.cleanup_queue import ( 

44 CleanupQueueOperation, 

45 CleanupQueueOperationResult, 

46 CleanupQueueReservation, 

47 WorkerCleanupQueue, 

48) 

49from prefect.types._datetime import now 

50 

51try: 

52 from prometheus_client import Counter as _Counter 

53except ImportError: # pragma: no cover 

54 _Counter = None # type: ignore[assignment,misc] 

55 

56if TYPE_CHECKING: 56 ↛ 57line 56 didn't jump to line 57 because the condition on line 56 was never true

57 from prefect.server.database.orm_models import WorkPool as ORMWorkPool 

58 

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

60 

61# --------------------------------------------------------------------------- 

62# Prometheus metrics 

63# --------------------------------------------------------------------------- 

64if _Counter is not None: 64 ↛ 90line 64 didn't jump to line 90 because the condition on line 64 was always true

65 WORKER_CHANNEL_CONNECTIONS = _Counter( 

66 "prefect_worker_channel_connections_total", 

67 "Worker channel connection lifecycle events.", 

68 ["event"], 

69 ) 

70 WORKER_CHANNEL_CLOSE_REASONS = _Counter( 

71 "prefect_worker_channel_close_reasons_total", 

72 "Worker channel close reasons.", 

73 ["reason"], 

74 ) 

75 WORKER_CHANNEL_SNAPSHOTS = _Counter( 

76 "prefect_worker_channel_snapshots_total", 

77 "Worker channel snapshot events.", 

78 ["event"], 

79 ) 

80 WORKER_CHANNEL_HEARTBEAT_FAILURES = _Counter( 

81 "prefect_worker_channel_heartbeat_failures_total", 

82 "Worker channel heartbeat persistence failures.", 

83 ) 

84 WORKER_CHANNEL_CLEANUP_DELIVERIES = _Counter( 

85 "prefect_worker_channel_cleanup_deliveries_total", 

86 "Worker channel cleanup message delivery outcomes.", 

87 ["status"], 

88 ) 

89else: 

90 WORKER_CHANNEL_CONNECTIONS = None 

91 WORKER_CHANNEL_CLOSE_REASONS = None 

92 WORKER_CHANNEL_SNAPSHOTS = None 

93 WORKER_CHANNEL_HEARTBEAT_FAILURES = None 

94 WORKER_CHANNEL_CLEANUP_DELIVERIES = None 

95 

96WORKER_CHANNEL_SNAPSHOT_TOPIC = "work-pool-worker-channel-snapshots" 

97_WORKER_CHANNEL_SNAPSHOT_BUFFER_SIZE = 1 

98_WORKER_CHANNEL_SNAPSHOT_COALESCE_SECONDS = 0.05 

99 

100WORK_POOL_FIELDS_THAT_TRIGGER_SNAPSHOTS = frozenset( 

101 { 

102 "base_job_template", 

103 "concurrency_limit", 

104 "is_paused", 

105 "storage_configuration", 

106 } 

107) 

108 

109 

110class WorkerChannelSnapshotInvalidation(PrefectBaseModel): 

111 work_pool_id: UUID 

112 reason: str 

113 work_pool_deleted: bool = False 

114 

115 def targets(self, *, work_pool_id: UUID) -> bool: 

116 return self.work_pool_id == work_pool_id 

117 

118 

119def work_pool_update_triggers_snapshot(update_values: Mapping[str, Any]) -> bool: 

120 return bool(WORK_POOL_FIELDS_THAT_TRIGGER_SNAPSHOTS.intersection(update_values)) 

121 

122 

123class WorkerChannelConnection: 

124 def __init__( 

125 self, 

126 *, 

127 websocket: WebSocket, 

128 db: PrefectDBInterface, 

129 work_pool_name: str, 

130 work_pool_id: UUID, 

131 consumer_id: UUID, 

132 worker_name: str, 

133 work_pool_updated: datetime, 

134 cleanup_queue: WorkerCleanupQueue | None = None, 

135 cleanup_kinds: tuple[CleanupKind, ...] = (), 

136 cleanup_work_queue_ids: tuple[UUID, ...] = (), 

137 max_cleanup_concurrency: int = 0, 

138 cleanup_registry: WorkerCleanupConnectionRegistry = ( 

139 WORKER_CLEANUP_CONNECTION_REGISTRY 

140 ), 

141 ) -> None: 

142 self.websocket = websocket 

143 self.db = db 

144 self.work_pool_name = work_pool_name 

145 self.work_pool_id = work_pool_id 

146 self.consumer_id = consumer_id 

147 self.worker_name = worker_name 

148 self._work_pool_updated = work_pool_updated 

149 self._next_snapshot_sequence = 2 

150 self._snapshot_queue: asyncio.Queue[WorkerChannelSnapshotInvalidation] = ( 

151 asyncio.Queue(maxsize=_WORKER_CHANNEL_SNAPSHOT_BUFFER_SIZE) 

152 ) 

153 self._send_lock = asyncio.Lock() 

154 self._closed = asyncio.Event() 

155 self._ready_sent = asyncio.Event() 

156 self._cleanup_queue = cleanup_queue 

157 self._cleanup_kinds = cleanup_kinds 

158 self._cleanup_work_queue_ids = cleanup_work_queue_ids 

159 self._max_cleanup_concurrency = max_cleanup_concurrency 

160 self._cleanup_registry = cleanup_registry 

161 

162 @property 

163 def cleanup_enabled(self) -> bool: 

164 return ( 

165 self._cleanup_queue is not None 

166 and self._cleanup_kinds 

167 and self._max_cleanup_concurrency > 0 

168 ) 

169 

170 async def run( 

171 self, ready: WorkerReadyFrame, consumer_kwargs: Mapping[str, Any] 

172 ) -> None: 

173 logger.debug( 

174 "Worker channel connection opened: " 

175 "work_pool=%s worker_name=%s cleanup_enabled=%s", 

176 self.work_pool_name, 

177 self.worker_name, 

178 self.cleanup_enabled, 

179 ) 

180 if WORKER_CHANNEL_CONNECTIONS is not None: 

181 WORKER_CHANNEL_CONNECTIONS.labels(event="opened").inc() 

182 

183 try: 

184 if self.cleanup_enabled: 

185 async with self._cleanup_registry.register(self): 

186 await self._run(ready, consumer_kwargs) 

187 return 

188 

189 await self._run(ready, consumer_kwargs) 

190 finally: 

191 logger.debug( 

192 "Worker channel connection closed: work_pool=%s worker_name=%s", 

193 self.work_pool_name, 

194 self.worker_name, 

195 ) 

196 if WORKER_CHANNEL_CONNECTIONS is not None: 

197 WORKER_CHANNEL_CONNECTIONS.labels(event="closed").inc() 

198 

199 async def _run( 

200 self, ready: WorkerReadyFrame, consumer_kwargs: Mapping[str, Any] 

201 ) -> None: 

202 send_task = asyncio.create_task(self._send_loop(ready)) 

203 receive_task = asyncio.create_task(self._receive_loop()) 

204 fanout_task = asyncio.create_task(self._fanout_loop(consumer_kwargs)) 

205 tasks = {send_task, receive_task, fanout_task} 

206 

207 try: 

208 done, pending = await asyncio.wait( 

209 tasks, return_when=asyncio.FIRST_COMPLETED 

210 ) 

211 except asyncio.CancelledError: 

212 self._closed.set() 

213 for task in tasks: 

214 task.cancel() 

215 await asyncio.gather(*tasks, return_exceptions=True) 

216 return 

217 except BaseException: 

218 self._closed.set() 

219 for task in tasks: 

220 task.cancel() 

221 await asyncio.gather(*tasks, return_exceptions=True) 

222 raise 

223 

224 self._closed.set() 

225 for task in pending: 

226 task.cancel() 

227 try: 

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

229 except asyncio.CancelledError: 

230 return 

231 

232 for task in done: 

233 if task.cancelled(): 

234 continue 

235 exception = task.exception() 

236 if exception is not None: 

237 raise exception 

238 

239 async def close(self, close_reason: WorkerChannelCloseReason) -> None: 

240 if self._closed.is_set(): 

241 return 

242 

243 self._closed.set() 

244 logger.debug( 

245 "Worker channel closing: work_pool=%s worker_name=%s reason=%s", 

246 self.work_pool_name, 

247 self.worker_name, 

248 close_reason.value, 

249 ) 

250 if WORKER_CHANNEL_CLOSE_REASONS is not None: 

251 WORKER_CHANNEL_CLOSE_REASONS.labels(reason=close_reason.value).inc() 

252 

253 async with self._send_lock: 

254 await close_worker_channel(self.websocket, close_reason) 

255 

256 def queue_snapshot( 

257 self, 

258 invalidation: WorkerChannelSnapshotInvalidation, 

259 ) -> None: 

260 if self._closed.is_set() or not invalidation.targets( 

261 work_pool_id=self.work_pool_id, 

262 ): 

263 return 

264 

265 while True: 

266 try: 

267 self._snapshot_queue.put_nowait(invalidation) 

268 return 

269 except asyncio.QueueFull: 

270 try: 

271 self._snapshot_queue.get_nowait() 

272 logger.debug( 

273 "Worker channel snapshot backpressure drop: " 

274 "work_pool=%s reason=%s", 

275 self.work_pool_name, 

276 invalidation.reason, 

277 ) 

278 if WORKER_CHANNEL_SNAPSHOTS is not None: 

279 WORKER_CHANNEL_SNAPSHOTS.labels( 

280 event="backpressure_dropped" 

281 ).inc() 

282 except asyncio.QueueEmpty: 

283 continue 

284 

285 async def has_cleanup_capacity(self) -> bool: 

286 if ( 

287 not self.cleanup_enabled 

288 or self._closed.is_set() 

289 or not self._ready_sent.is_set() 

290 ): 

291 return False 

292 

293 return await self._cleanup_registry.has_cleanup_capacity( 

294 self, self._max_cleanup_concurrency 

295 ) 

296 

297 async def dispatch_one_cleanup_message( 

298 self, 

299 *, 

300 cleanup_queue: WorkerCleanupQueue, 

301 allow_fallback_to_any_queue: bool, 

302 ) -> bool: 

303 if not await self.has_cleanup_capacity(): 

304 return False 

305 

306 preferred_work_queue_ids: tuple[UUID, ...] | None 

307 if allow_fallback_to_any_queue: 

308 preferred_work_queue_ids = None 

309 else: 

310 if not self._cleanup_work_queue_ids: 

311 return False 

312 preferred_work_queue_ids = self._cleanup_work_queue_ids 

313 

314 reservation = await cleanup_queue.reserve( 

315 work_pool_id=self.work_pool_id, 

316 cleanup_kinds=self._cleanup_kinds, 

317 preferred_work_queue_ids=preferred_work_queue_ids, 

318 allow_fallback_to_any_queue=allow_fallback_to_any_queue, 

319 ) 

320 if reservation is None: 

321 return False 

322 

323 if self._closed.is_set() or not self._ready_sent.is_set(): 

324 should_release = True 

325 else: 

326 should_release = not await self._cleanup_registry.track_cleanup_reservation( 

327 self, 

328 WorkerCleanupInFlight( 

329 message_id=reservation.message_id, 

330 reservation_token=reservation.reservation_token, 

331 lease_expires_at=reservation.lease_expires_at, 

332 ), 

333 self._max_cleanup_concurrency, 

334 ) 

335 

336 if should_release: 

337 logger.debug( 

338 "Worker channel cleanup delivery released before send: " 

339 "work_pool=%s message_id=%s reason=connection_unavailable", 

340 self.work_pool_name, 

341 reservation.message_id, 

342 ) 

343 if WORKER_CHANNEL_CLEANUP_DELIVERIES is not None: 

344 WORKER_CHANNEL_CLEANUP_DELIVERIES.labels(status="released").inc() 

345 await cleanup_queue.release( 

346 work_pool_id=self.work_pool_id, 

347 message_id=reservation.message_id, 

348 reservation_token=reservation.reservation_token, 

349 reason="connection_unavailable", 

350 ) 

351 return False 

352 

353 try: 

354 await self._send_frame(_build_cleanup_message_frame(reservation)) 

355 except asyncio.CancelledError: 

356 if WORKER_CHANNEL_CLEANUP_DELIVERIES is not None: 

357 WORKER_CHANNEL_CLEANUP_DELIVERIES.labels(status="failed").inc() 

358 await self._release_cleanup_delivery_failure( 

359 cleanup_queue=cleanup_queue, 

360 reservation=reservation, 

361 ) 

362 raise 

363 except subscriptions.NORMAL_DISCONNECT_EXCEPTIONS: 

364 self._closed.set() 

365 if WORKER_CHANNEL_CLEANUP_DELIVERIES is not None: 

366 WORKER_CHANNEL_CLEANUP_DELIVERIES.labels(status="failed").inc() 

367 await self._release_cleanup_delivery_failure( 

368 cleanup_queue=cleanup_queue, 

369 reservation=reservation, 

370 ) 

371 return False 

372 except Exception: 

373 logger.exception( 

374 "Worker channel cleanup message delivery failed: " 

375 "work_pool=%s message_id=%s cleanup_kind=%s", 

376 self.work_pool_name, 

377 reservation.message_id, 

378 reservation.kind, 

379 ) 

380 if WORKER_CHANNEL_CLEANUP_DELIVERIES is not None: 

381 WORKER_CHANNEL_CLEANUP_DELIVERIES.labels(status="failed").inc() 

382 await self._release_cleanup_delivery_failure( 

383 cleanup_queue=cleanup_queue, 

384 reservation=reservation, 

385 ) 

386 await self.close(WorkerChannelCloseReason.TRANSIENT_SERVER_ERROR) 

387 return False 

388 

389 logger.debug( 

390 "Worker channel cleanup message delivered: " 

391 "work_pool=%s message_id=%s cleanup_kind=%s " 

392 "delivery_count=%s", 

393 self.work_pool_name, 

394 reservation.message_id, 

395 reservation.kind, 

396 reservation.delivery_count, 

397 ) 

398 if WORKER_CHANNEL_CLEANUP_DELIVERIES is not None: 

399 WORKER_CHANNEL_CLEANUP_DELIVERIES.labels(status="delivered").inc() 

400 return True 

401 

402 async def _fanout_loop(self, consumer_kwargs: Mapping[str, Any]) -> None: 

403 if "subscription" in consumer_kwargs: 

404 consumer_kwargs = {**consumer_kwargs, "concurrency": 1} 

405 consumer = messaging.create_consumer(**consumer_kwargs) 

406 

407 async def handle_message(worker_channel_message: messaging.Message) -> None: 

408 invalidation = parse_snapshot_invalidation(worker_channel_message) 

409 self.queue_snapshot(invalidation) 

410 

411 await consumer.run(handle_message) 

412 

413 async def _send_loop(self, ready: WorkerReadyFrame) -> None: 

414 await self._send_frame(ready) 

415 self._ready_sent.set() 

416 if self.cleanup_enabled and self._cleanup_queue is not None: 

417 await self._dispatch_cleanup_available(self._cleanup_queue) 

418 self._cleanup_registry.wake_dispatcher(self.work_pool_id) 

419 

420 while not self._closed.is_set(): 

421 invalidation = await self._snapshot_queue.get() 

422 invalidation = await self._coalesce_snapshot_invalidations(invalidation) 

423 

424 if invalidation.work_pool_deleted: 

425 await self.close(WorkerChannelCloseReason.AUTHORIZATION_FAILED) 

426 return 

427 

428 frame = await self._build_snapshot_frame(invalidation) 

429 if frame is None: 

430 await self.close(WorkerChannelCloseReason.AUTHORIZATION_FAILED) 

431 return 

432 

433 await self._send_frame(frame) 

434 logger.debug( 

435 "Worker channel snapshot sent: work_pool=%s sequence=%s reason=%s", 

436 self.work_pool_name, 

437 frame.payload.snapshot_sequence, 

438 invalidation.reason, 

439 ) 

440 if WORKER_CHANNEL_SNAPSHOTS is not None: 

441 WORKER_CHANNEL_SNAPSHOTS.labels(event="sent").inc() 

442 

443 async def _send_frame( 

444 self, 

445 frame: ( 

446 WorkerReadyFrame 

447 | WorkPoolSnapshotFrame 

448 | CleanupMessageFrame 

449 | CleanupOperationResultFrame 

450 ), 

451 ) -> None: 

452 async with self._send_lock: 

453 await self.websocket.send_json(frame.model_dump(mode="json")) 

454 

455 async def _release_cleanup_delivery_failure( 

456 self, 

457 *, 

458 cleanup_queue: WorkerCleanupQueue, 

459 reservation: CleanupQueueReservation, 

460 ) -> None: 

461 await self._forget_cleanup_reservation( 

462 message_id=reservation.message_id, 

463 reservation_token=reservation.reservation_token, 

464 ) 

465 await cleanup_queue.release( 

466 work_pool_id=self.work_pool_id, 

467 message_id=reservation.message_id, 

468 reservation_token=reservation.reservation_token, 

469 reason="delivery_failed", 

470 ) 

471 

472 async def _handle_cleanup_operation( 

473 self, 

474 frame: CleanupAckFrame | CleanupReleaseFrame | CleanupRenewFrame, 

475 ) -> None: 

476 operation = _cleanup_operation_from_frame(frame) 

477 if not self.cleanup_enabled or self._cleanup_queue is None: 

478 await self._send_frame( 

479 _build_cleanup_operation_result_frame( 

480 request_frame_id=frame.id, 

481 result=CleanupQueueOperationResult( 

482 message_id=frame.payload.message_id, 

483 operation=operation, 

484 status="unauthorized", 

485 reason="cleanup_delivery_not_accepted", 

486 ), 

487 ) 

488 ) 

489 return 

490 

491 cleanup_queue = self._cleanup_queue 

492 try: 

493 if isinstance(frame, CleanupAckFrame): 

494 result = await cleanup_queue.ack( 

495 work_pool_id=self.work_pool_id, 

496 message_id=frame.payload.message_id, 

497 reservation_token=frame.payload.reservation_token, 

498 ) 

499 elif isinstance(frame, CleanupReleaseFrame): 

500 result = await cleanup_queue.release( 

501 work_pool_id=self.work_pool_id, 

502 message_id=frame.payload.message_id, 

503 reservation_token=frame.payload.reservation_token, 

504 reason=frame.payload.reason, 

505 ) 

506 else: 

507 result = await cleanup_queue.renew( 

508 work_pool_id=self.work_pool_id, 

509 message_id=frame.payload.message_id, 

510 reservation_token=frame.payload.reservation_token, 

511 ) 

512 except asyncio.CancelledError: 

513 raise 

514 except Exception: 

515 logger.exception( 

516 "Worker cleanup queue operation failed: " 

517 "work_pool=%s operation=%s message_id=%s", 

518 self.work_pool_name, 

519 operation, 

520 frame.payload.message_id, 

521 ) 

522 result = CleanupQueueOperationResult( 

523 message_id=frame.payload.message_id, 

524 operation=operation, 

525 status="error", 

526 reason="cleanup_queue_operation_failed", 

527 ) 

528 

529 synced_before_send = result.operation == "renew" and result.status == "accepted" 

530 if synced_before_send: 

531 await self._sync_cleanup_operation_result( 

532 reservation_token=frame.payload.reservation_token, 

533 result=result, 

534 ) 

535 

536 send_succeeded = False 

537 try: 

538 await self._send_frame( 

539 _build_cleanup_operation_result_frame( 

540 request_frame_id=frame.id, 

541 result=result, 

542 ) 

543 ) 

544 send_succeeded = True 

545 except subscriptions.NORMAL_DISCONNECT_EXCEPTIONS: 

546 self._closed.set() 

547 raise 

548 finally: 

549 if not synced_before_send: 

550 freed_capacity = await self._sync_cleanup_operation_result( 

551 reservation_token=frame.payload.reservation_token, 

552 result=result, 

553 ) 

554 if send_succeeded and freed_capacity: 

555 await self._dispatch_cleanup_available(cleanup_queue) 

556 

557 async def _dispatch_cleanup_available( 

558 self, 

559 cleanup_queue: WorkerCleanupQueue, 

560 ) -> None: 

561 try: 

562 await self._cleanup_registry.dispatch_available( 

563 work_pool_id=self.work_pool_id, 

564 cleanup_queue=cleanup_queue, 

565 ) 

566 except asyncio.CancelledError: 

567 raise 

568 except Exception: 

569 logger.exception("Worker channel cleanup dispatch failed") 

570 self._cleanup_registry.wake_dispatcher(self.work_pool_id) 

571 

572 async def _sync_cleanup_operation_result( 

573 self, 

574 *, 

575 reservation_token: str, 

576 result: CleanupQueueOperationResult, 

577 ) -> bool: 

578 if result.status == "error": 

579 return False 

580 

581 if result.operation == "renew" and result.status == "accepted": 

582 if result.lease_expires_at is not None: 

583 await self._cleanup_registry.update_cleanup_lease( 

584 self, 

585 reservation_token=reservation_token, 

586 lease_expires_at=result.lease_expires_at, 

587 ) 

588 return False 

589 

590 return await self._forget_cleanup_reservation( 

591 message_id=result.message_id, 

592 reservation_token=reservation_token, 

593 ) 

594 

595 async def _forget_cleanup_reservation( 

596 self, *, message_id: UUID, reservation_token: str 

597 ) -> bool: 

598 return await self._cleanup_registry.forget_cleanup_reservation( 

599 self, 

600 message_id=message_id, 

601 reservation_token=reservation_token, 

602 ) 

603 

604 async def _coalesce_snapshot_invalidations( 

605 self, 

606 invalidation: WorkerChannelSnapshotInvalidation, 

607 ) -> WorkerChannelSnapshotInvalidation: 

608 await asyncio.sleep(_WORKER_CHANNEL_SNAPSHOT_COALESCE_SECONDS) 

609 

610 coalesced_count = 0 

611 while True: 

612 try: 

613 invalidation = self._snapshot_queue.get_nowait() 

614 coalesced_count += 1 

615 except asyncio.QueueEmpty: 

616 if coalesced_count > 0: 

617 logger.debug( 

618 "Worker channel snapshots coalesced: " 

619 "work_pool=%s coalesced_count=%s", 

620 self.work_pool_name, 

621 coalesced_count, 

622 ) 

623 if WORKER_CHANNEL_SNAPSHOTS is not None: 

624 WORKER_CHANNEL_SNAPSHOTS.labels(event="coalesced").inc( 

625 coalesced_count 

626 ) 

627 return invalidation 

628 

629 async def _build_snapshot_frame( 

630 self, 

631 invalidation: WorkerChannelSnapshotInvalidation, 

632 ) -> WorkPoolSnapshotFrame | None: 

633 async with self.db.session_context() as session: 

634 work_pool = await models.workers.read_work_pool( 

635 session=session, 

636 work_pool_id=invalidation.work_pool_id, 

637 ) 

638 if work_pool is None: 

639 return None 

640 

641 self._work_pool_updated = work_pool.updated 

642 

643 payload = WorkPoolSnapshotPayload( 

644 snapshot_sequence=self._next_snapshot_sequence, 

645 reason=invalidation.reason, 

646 work_pool=await build_worker_channel_work_pool_snapshot( 

647 session=session, 

648 work_pool=work_pool, 

649 ), 

650 ) 

651 

652 self._next_snapshot_sequence += 1 

653 return WorkPoolSnapshotFrame( 

654 type="work_pool.snapshot.v1", 

655 id=uuid7(), 

656 sent_at=now("UTC"), 

657 payload=payload, 

658 ) 

659 

660 async def _receive_loop(self) -> None: 

661 while not self._closed.is_set(): 

662 try: 

663 message = await self.websocket.receive_json() 

664 frame = validate_worker_channel_frame(message) 

665 except subscriptions.NORMAL_DISCONNECT_EXCEPTIONS: 

666 return 

667 except ValidationError: 

668 await self.close(WorkerChannelCloseReason.PROTOCOL_ERROR) 

669 return 

670 except ValueError: 

671 await self.close(WorkerChannelCloseReason.PROTOCOL_ERROR) 

672 return 

673 

674 if isinstance(frame, WorkerHeartbeatFrame): 

675 if ( 

676 frame.payload.consumer_id != self.consumer_id 

677 or frame.payload.worker_name != self.worker_name 

678 ): 

679 await self.close(WorkerChannelCloseReason.PROTOCOL_ERROR) 

680 return 

681 

682 try: 

683 async with self.db.session_context( 

684 begin_transaction=True 

685 ) as session: 

686 pool_updated = await _persist_worker_channel_heartbeat( 

687 session=session, 

688 work_pool_name=self.work_pool_name, 

689 frame=frame, 

690 ) 

691 except Exception: 

692 logger.exception( 

693 "Worker channel heartbeat persistence failed: " 

694 "work_pool=%s worker_name=%s", 

695 self.work_pool_name, 

696 self.worker_name, 

697 ) 

698 if WORKER_CHANNEL_HEARTBEAT_FAILURES is not None: 

699 WORKER_CHANNEL_HEARTBEAT_FAILURES.inc() 

700 await self.close( 

701 WorkerChannelCloseReason.HEARTBEAT_PERSISTENCE_FAILED 

702 ) 

703 return 

704 

705 if pool_updated != self._work_pool_updated: 

706 self.queue_snapshot( 

707 WorkerChannelSnapshotInvalidation( 

708 work_pool_id=self.work_pool_id, 

709 reason="heartbeat_reconciliation", 

710 ) 

711 ) 

712 continue 

713 

714 if isinstance( 

715 frame, (CleanupAckFrame, CleanupReleaseFrame, CleanupRenewFrame) 

716 ): 

717 await self._handle_cleanup_operation(frame) 

718 continue 

719 

720 await self.close(WorkerChannelCloseReason.PROTOCOL_ERROR) 

721 return 

722 

723 

724async def close_worker_channel( 

725 websocket: WebSocket, close_reason: WorkerChannelCloseReason 

726) -> None: 

727 policy = WORKER_CHANNEL_CLOSE_POLICIES[close_reason] 

728 await websocket.close(code=policy.websocket_code, reason=close_reason.value) 

729 

730 

731async def build_worker_channel_work_pool_snapshot( 

732 session: AsyncSession, 

733 work_pool: ORMWorkPool, 

734) -> WorkPoolSnapshot: 

735 work_pool_response = schemas.responses.WorkPoolResponse.model_validate( 

736 work_pool, from_attributes=True 

737 ) 

738 

739 if work_pool_response.concurrency_limit is not None: 

740 work_pool_response.active_slots = ( 

741 await models.workers.count_work_pool_active_slots( 

742 session=session, 

743 work_pool_id=work_pool.id, 

744 ) 

745 ) 

746 

747 return WorkPoolSnapshot.model_validate(work_pool_response.model_dump(mode="json")) 

748 

749 

750def _build_cleanup_message_frame( 

751 reservation: CleanupQueueReservation, 

752) -> CleanupMessageFrame: 

753 return CleanupMessageFrame( 

754 type="cleanup.message.v1", 

755 id=uuid7(), 

756 sent_at=now("UTC"), 

757 payload={ 

758 "message_id": reservation.message_id, 

759 "kind": reservation.kind, 

760 "reservation_token": reservation.reservation_token, 

761 "lease_expires_at": reservation.lease_expires_at, 

762 "delivery_count": reservation.delivery_count, 

763 "work_queue_id": reservation.work_queue_id, 

764 "target": reservation.target, 

765 "data": reservation.data, 

766 }, 

767 ) 

768 

769 

770def _build_cleanup_operation_result_frame( 

771 *, 

772 request_frame_id: UUID, 

773 result: CleanupQueueOperationResult, 

774) -> CleanupOperationResultFrame: 

775 return CleanupOperationResultFrame( 

776 type="cleanup.operation_result.v1", 

777 id=uuid7(), 

778 sent_at=now("UTC"), 

779 payload={ 

780 "request_frame_id": request_frame_id, 

781 "message_id": result.message_id, 

782 "operation": result.operation, 

783 "status": result.status, 

784 "lease_expires_at": result.lease_expires_at, 

785 "reason": result.reason, 

786 "detail": None, 

787 }, 

788 ) 

789 

790 

791def _cleanup_operation_from_frame( 

792 frame: CleanupAckFrame | CleanupReleaseFrame | CleanupRenewFrame, 

793) -> CleanupQueueOperation: 

794 if isinstance(frame, CleanupAckFrame): 

795 return "ack" 

796 if isinstance(frame, CleanupReleaseFrame): 

797 return "release" 

798 return "renew" 

799 

800 

801async def _persist_worker_channel_heartbeat( 

802 session: AsyncSession, 

803 work_pool_name: str, 

804 frame: WorkerHeartbeatFrame, 

805) -> datetime: 

806 work_pool = await models.workers.read_work_pool_by_name( 

807 session=session, 

808 work_pool_name=work_pool_name, 

809 ) 

810 if work_pool is None: 

811 raise RuntimeError("Worker channel work pool no longer exists") 

812 

813 await models.workers.record_worker_heartbeat( 

814 session=session, 

815 work_pool=work_pool, 

816 worker_name=frame.payload.worker_name, 

817 heartbeat_interval_seconds=frame.payload.heartbeat_interval_seconds, 

818 emit_status_change=emit_work_pool_status_event, 

819 ) 

820 

821 return work_pool.updated 

822 

823 

824async def publish_snapshot_invalidation( 

825 invalidation: WorkerChannelSnapshotInvalidation, 

826) -> None: 

827 async with messaging.create_publisher( 

828 topic=WORKER_CHANNEL_SNAPSHOT_TOPIC 

829 ) as publisher: 

830 await publisher.publish_data( 

831 invalidation.model_dump_json().encode(), 

832 attributes={ 

833 "work_pool_id": str(invalidation.work_pool_id), 

834 "reason": invalidation.reason, 

835 }, 

836 ) 

837 

838 

839def parse_snapshot_invalidation( 

840 message: messaging.Message, 

841) -> WorkerChannelSnapshotInvalidation: 

842 data = message.data.encode() if isinstance(message.data, str) else message.data 

843 return WorkerChannelSnapshotInvalidation.model_validate_json(data)