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
« 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.abc import Mapping
6from datetime import datetime
7from typing import TYPE_CHECKING, Any
8from uuid import UUID
10from pydantic import ValidationError
11from sqlalchemy.ext.asyncio import AsyncSession
12from starlette.websockets import WebSocket
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
51try:
52 from prometheus_client import Counter as _Counter
53except ImportError: # pragma: no cover
54 _Counter = None # type: ignore[assignment,misc]
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
59logger: logging.Logger = get_logger("prefect.server.utilities.worker_channel")
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
96WORKER_CHANNEL_SNAPSHOT_TOPIC = "work-pool-worker-channel-snapshots"
97_WORKER_CHANNEL_SNAPSHOT_BUFFER_SIZE = 1
98_WORKER_CHANNEL_SNAPSHOT_COALESCE_SECONDS = 0.05
100WORK_POOL_FIELDS_THAT_TRIGGER_SNAPSHOTS = frozenset(
101 {
102 "base_job_template",
103 "concurrency_limit",
104 "is_paused",
105 "storage_configuration",
106 }
107)
110class WorkerChannelSnapshotInvalidation(PrefectBaseModel):
111 work_pool_id: UUID
112 reason: str
113 work_pool_deleted: bool = False
115 def targets(self, *, work_pool_id: UUID) -> bool:
116 return self.work_pool_id == work_pool_id
119def work_pool_update_triggers_snapshot(update_values: Mapping[str, Any]) -> bool:
120 return bool(WORK_POOL_FIELDS_THAT_TRIGGER_SNAPSHOTS.intersection(update_values))
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
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 )
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()
183 try:
184 if self.cleanup_enabled:
185 async with self._cleanup_registry.register(self):
186 await self._run(ready, consumer_kwargs)
187 return
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()
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}
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
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
232 for task in done:
233 if task.cancelled():
234 continue
235 exception = task.exception()
236 if exception is not None:
237 raise exception
239 async def close(self, close_reason: WorkerChannelCloseReason) -> None:
240 if self._closed.is_set():
241 return
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()
253 async with self._send_lock:
254 await close_worker_channel(self.websocket, close_reason)
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
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
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
293 return await self._cleanup_registry.has_cleanup_capacity(
294 self, self._max_cleanup_concurrency
295 )
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
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
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
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 )
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
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
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
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)
407 async def handle_message(worker_channel_message: messaging.Message) -> None:
408 invalidation = parse_snapshot_invalidation(worker_channel_message)
409 self.queue_snapshot(invalidation)
411 await consumer.run(handle_message)
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)
420 while not self._closed.is_set():
421 invalidation = await self._snapshot_queue.get()
422 invalidation = await self._coalesce_snapshot_invalidations(invalidation)
424 if invalidation.work_pool_deleted:
425 await self.close(WorkerChannelCloseReason.AUTHORIZATION_FAILED)
426 return
428 frame = await self._build_snapshot_frame(invalidation)
429 if frame is None:
430 await self.close(WorkerChannelCloseReason.AUTHORIZATION_FAILED)
431 return
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()
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"))
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 )
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
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 )
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 )
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)
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)
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
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
590 return await self._forget_cleanup_reservation(
591 message_id=result.message_id,
592 reservation_token=reservation_token,
593 )
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 )
604 async def _coalesce_snapshot_invalidations(
605 self,
606 invalidation: WorkerChannelSnapshotInvalidation,
607 ) -> WorkerChannelSnapshotInvalidation:
608 await asyncio.sleep(_WORKER_CHANNEL_SNAPSHOT_COALESCE_SECONDS)
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
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
641 self._work_pool_updated = work_pool.updated
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 )
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 )
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
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
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
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
714 if isinstance(
715 frame, (CleanupAckFrame, CleanupReleaseFrame, CleanupRenewFrame)
716 ):
717 await self._handle_cleanup_operation(frame)
718 continue
720 await self.close(WorkerChannelCloseReason.PROTOCOL_ERROR)
721 return
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)
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 )
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 )
747 return WorkPoolSnapshot.model_validate(work_pool_response.model_dump(mode="json"))
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 )
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 )
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"
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")
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 )
821 return work_pool.updated
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 )
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)