Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/models/workers.py: 51%
381 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
1"""
2Functions for interacting with worker ORM objects.
3Intended for internal use by the Prefect REST API.
4"""
6import datetime
7from typing import (
8 Any,
9 Awaitable,
10 Callable,
11 Dict,
12 List,
13 Optional,
14 Sequence,
15 Union,
16)
17from uuid import UUID
19import sqlalchemy as sa
20from sqlalchemy import delete, select
21from sqlalchemy.ext.asyncio import AsyncSession
23import prefect.server.schemas as schemas
24from prefect._internal.uuid7 import uuid7
25from prefect.server.database import PrefectDBInterface, db_injector, orm_models
26from prefect.server.events import clients
27from prefect.server.events.clients import PrefectServerEventsClient
28from prefect.server.exceptions import ObjectNotFoundError
29from prefect.server.models.events import (
30 work_pool_created_event,
31 work_pool_deleted_event,
32 work_pool_status_event,
33 work_pool_updated_event, # Add this
34 work_queue_created_event,
35 work_queue_deleted_event,
36)
37from prefect.server.schemas.statuses import WorkQueueStatus
38from prefect.server.utilities.database import UUID as PrefectUUID
39from prefect.types._datetime import DateTime, now
41DEFAULT_AGENT_WORK_POOL_NAME = "default-agent-pool"
43# -----------------------------------------------------
44# --
45# --
46# -- Work Pools
47# --
48# --
49# -----------------------------------------------------
52@db_injector
53async def create_work_pool(
54 db: PrefectDBInterface,
55 session: AsyncSession,
56 work_pool: Union[schemas.core.WorkPool, schemas.actions.WorkPoolCreate],
57) -> orm_models.WorkPool:
58 """
59 Creates a work pool.
61 If a WorkPool with the same name exists, an error will be thrown.
63 Args:
64 session (AsyncSession): a database session
65 work_pool (schemas.core.WorkPool): a WorkPool model
67 Returns:
68 orm_models.WorkPool: the newly-created WorkPool
70 """
72 pool = db.WorkPool(**work_pool.model_dump(exclude={"active_slots"}))
74 if pool.type != "prefect-agent":
75 if pool.is_paused:
76 pool.status = schemas.statuses.WorkPoolStatus.PAUSED
77 else:
78 pool.status = schemas.statuses.WorkPoolStatus.NOT_READY
80 session.add(pool)
81 await session.flush()
83 default_queue = await create_work_queue(
84 session=session,
85 work_pool_id=pool.id,
86 work_queue=schemas.actions.WorkQueueCreate(
87 name="default", description="The work pool's default queue."
88 ),
89 )
91 pool.default_queue_id = default_queue.id # type: ignore
92 await session.flush()
94 await emit_work_pool_created_event(pool)
96 return pool
99@db_injector
100async def read_work_pool(
101 db: PrefectDBInterface, session: AsyncSession, work_pool_id: UUID
102) -> Optional[orm_models.WorkPool]:
103 """
104 Reads a WorkPool by id.
106 Args:
107 session (AsyncSession): A database session
108 work_pool_id (UUID): a WorkPool id
110 Returns:
111 orm_models.WorkPool: the WorkPool
112 """
113 query = sa.select(db.WorkPool).where(db.WorkPool.id == work_pool_id).limit(1)
114 result = await session.execute(query)
115 return result.scalar()
118@db_injector
119async def read_work_pool_by_name(
120 db: PrefectDBInterface, session: AsyncSession, work_pool_name: str
121) -> Optional[orm_models.WorkPool]:
122 """
123 Reads a WorkPool by name.
125 Args:
126 session (AsyncSession): A database session
127 work_pool_name (str): a WorkPool name
129 Returns:
130 orm_models.WorkPool: the WorkPool
131 """
132 query = sa.select(db.WorkPool).where(db.WorkPool.name == work_pool_name).limit(1)
133 result = await session.execute(query)
134 return result.scalar()
137@db_injector
138async def read_work_pools(
139 db: PrefectDBInterface,
140 session: AsyncSession,
141 work_pool_filter: Optional[schemas.filters.WorkPoolFilter] = None,
142 offset: Optional[int] = None,
143 limit: Optional[int] = None,
144) -> Sequence[orm_models.WorkPool]:
145 """
146 Read worker configs.
148 Args:
149 session: A database session
150 offset: Query offset
151 limit: Query limit
152 Returns:
153 List[orm_models.WorkPool]: worker configs
154 """
156 query = select(db.WorkPool).order_by(db.WorkPool.name)
158 if work_pool_filter is not None:
159 query = query.where(work_pool_filter.as_sql_filter())
160 if offset is not None: 160 ↛ 162line 160 didn't jump to line 162 because the condition on line 160 was always true
161 query = query.offset(offset)
162 if limit is not None: 162 ↛ 165line 162 didn't jump to line 165 because the condition on line 162 was always true
163 query = query.limit(limit)
165 result = await session.execute(query)
166 return result.scalars().unique().all()
169@db_injector
170async def count_work_pools(
171 db: PrefectDBInterface,
172 session: AsyncSession,
173 work_pool_filter: Optional[schemas.filters.WorkPoolFilter] = None,
174) -> int:
175 """
176 Read worker configs.
178 Args:
179 session: A database session
180 work_pool_filter: filter criteria to apply to the count
181 Returns:
182 int: the count of work pools matching the criteria
183 """
185 query = select(sa.func.count()).select_from(db.WorkPool)
187 if work_pool_filter is not None:
188 query = query.where(work_pool_filter.as_sql_filter())
190 result = await session.execute(query)
191 return result.scalar_one()
194# States counted against both pool-level and queue-level concurrency by the
195# worker scheduling SQL templates (get-runs-from-worker-queues.sql.jinja).
196SLOT_OCCUPYING_STATES = {
197 schemas.states.StateType.PENDING,
198 schemas.states.StateType.RUNNING,
199}
202@db_injector
203async def count_work_pool_active_slots(
204 db: PrefectDBInterface,
205 session: AsyncSession,
206 work_pool_id: UUID,
207) -> int:
208 """
209 Count flow runs in slot-occupying states (Pending, Running) for a given
210 work pool. Does not filter on queue pause status — paused queues may
211 still have running/pending runs consuming resources. This matches the
212 behavior of count_work_pool_slot_holders / get_work_pool_slot_holders.
213 """
214 query = (
215 select(sa.func.count())
216 .select_from(db.FlowRun)
217 .join(db.WorkQueue, db.FlowRun.work_queue_id == db.WorkQueue.id)
218 .where(
219 db.WorkQueue.work_pool_id == work_pool_id,
220 db.FlowRun.state_type.in_(SLOT_OCCUPYING_STATES),
221 )
222 )
223 result = await session.execute(query)
224 return result.scalar_one()
227@db_injector
228async def count_work_pool_active_slots_bulk(
229 db: PrefectDBInterface,
230 session: AsyncSession,
231 work_pool_ids: Sequence[UUID],
232) -> dict[UUID, int]:
233 """
234 Count active slots for multiple work pools in a single query.
235 Returns a mapping of work_pool_id -> active slot count.
236 Does not filter on queue pause status (see count_work_pool_active_slots).
237 """
238 if not work_pool_ids: 238 ↛ 239line 238 didn't jump to line 239 because the condition on line 238 was never true
239 return {}
241 query = (
242 select(
243 db.WorkQueue.work_pool_id,
244 sa.func.count(db.FlowRun.id),
245 )
246 .select_from(db.FlowRun)
247 .join(db.WorkQueue, db.FlowRun.work_queue_id == db.WorkQueue.id)
248 .where(
249 db.WorkQueue.work_pool_id.in_(work_pool_ids),
250 db.FlowRun.state_type.in_(SLOT_OCCUPYING_STATES),
251 )
252 .group_by(db.WorkQueue.work_pool_id)
253 )
254 result = await session.execute(query)
255 return dict(result.all())
258@db_injector
259async def count_work_queue_active_slots(
260 db: PrefectDBInterface,
261 session: AsyncSession,
262 work_queue_id: UUID,
263) -> int:
264 """
265 Count flow runs in slot-occupying states (Pending, Running) for a given
266 work queue under a work pool. Counts by work_queue_id FK only.
267 """
268 query = (
269 select(sa.func.count())
270 .select_from(db.FlowRun)
271 .where(
272 db.FlowRun.work_queue_id == work_queue_id,
273 db.FlowRun.state_type.in_(SLOT_OCCUPYING_STATES),
274 )
275 )
276 result = await session.execute(query)
277 return result.scalar_one()
280@db_injector
281async def count_work_queue_active_slots_bulk(
282 db: PrefectDBInterface,
283 session: AsyncSession,
284 work_queue_ids: Sequence[UUID],
285) -> dict[UUID, int]:
286 """
287 Count active slots for multiple work queues in a single query.
288 Returns a mapping of work_queue_id -> active slot count.
289 """
290 if not work_queue_ids: 290 ↛ 291line 290 didn't jump to line 291 because the condition on line 290 was never true
291 return {}
292 query = (
293 select(
294 db.FlowRun.work_queue_id,
295 sa.func.count(db.FlowRun.id),
296 )
297 .select_from(db.FlowRun)
298 .where(
299 db.FlowRun.work_queue_id.in_(work_queue_ids),
300 db.FlowRun.state_type.in_(SLOT_OCCUPYING_STATES),
301 )
302 .group_by(db.FlowRun.work_queue_id)
303 )
304 result = await session.execute(query)
305 return dict(result.all())
308@db_injector
309async def update_work_pool(
310 db: PrefectDBInterface,
311 session: AsyncSession,
312 work_pool_id: UUID,
313 work_pool: schemas.actions.WorkPoolUpdate,
314 emit_status_change: Optional[
315 Callable[
316 [UUID, DateTime, orm_models.WorkPool, orm_models.WorkPool],
317 Awaitable[None],
318 ]
319 ] = None,
320 emit_update_event: bool = True,
321) -> bool:
322 """
323 Update a WorkPool by id.
325 Args:
326 session (AsyncSession): A database session
327 work_pool_id (UUID): a WorkPool id
328 work_pool: the work pool data
329 emit_status_change: function to call when work pool
330 status is changed
331 emit_update_event: whether to emit an event for updated non-status fields
333 Returns:
334 bool: whether or not the worker was updated
335 """
336 # exclude_unset=True allows us to only update values provided by
337 # the user, ignoring any defaults on the model
338 update_data = work_pool.model_dump_for_orm(exclude_unset=True)
340 current_work_pool = await read_work_pool(session=session, work_pool_id=work_pool_id)
341 if not current_work_pool:
342 raise ObjectNotFoundError
344 # Remove this from the session so we have a copy of the current state before we
345 # update it; this will give us something to compare against when emitting events
346 session.expunge(current_work_pool)
348 if current_work_pool.type != "prefect-agent":
349 if update_data.get("is_paused"):
350 update_data["status"] = schemas.statuses.WorkPoolStatus.PAUSED
352 if update_data.get("is_paused") is False:
353 # If the work pool has any online workers, set the status to READY
354 # Otherwise set it to, NOT_READY
355 workers = await read_workers(
356 session=session,
357 work_pool_id=work_pool_id,
358 worker_filter=schemas.filters.WorkerFilter(
359 status=schemas.filters.WorkerFilterStatus(
360 any_=[schemas.statuses.WorkerStatus.ONLINE]
361 )
362 ),
363 )
364 if len(workers) > 0:
365 update_data["status"] = schemas.statuses.WorkPoolStatus.READY
366 else:
367 update_data["status"] = schemas.statuses.WorkPoolStatus.NOT_READY
369 if "status" in update_data:
370 update_data["last_status_event_id"] = uuid7()
371 update_data["last_transitioned_status_at"] = now("UTC")
373 update_stmt = (
374 sa.update(db.WorkPool)
375 .where(db.WorkPool.id == work_pool_id)
376 .values(**update_data)
377 )
378 result = await session.execute(update_stmt)
380 updated = result.rowcount > 0
381 if updated:
382 wp = await read_work_pool(session=session, work_pool_id=work_pool_id)
384 assert wp is not None
385 assert current_work_pool is not wp
387 # Detect which fields actually changed (excluding status and internal fields)
388 # Fields that should trigger update events (user-updatable fields)
389 WORK_POOL_EVENT_FIELDS = {
390 # "is_paused", # Handled with status
391 "description",
392 "base_job_template",
393 "concurrency_limit",
394 "storage_configuration",
395 }
397 changed_fields = {}
398 for field in update_data.keys():
399 if field not in WORK_POOL_EVENT_FIELDS or field == "status":
400 continue
402 old_value = getattr(current_work_pool, field, None)
403 new_value = getattr(wp, field, None)
405 # Compare values (handle different types)
406 if old_value != new_value:
407 changed_fields[field] = {
408 "from": old_value,
409 "to": new_value,
410 }
412 # Emit event for non-status field changes
413 if changed_fields and emit_update_event:
414 await emit_work_pool_updated_event(
415 session=session,
416 work_pool=wp,
417 changed_fields=changed_fields,
418 )
420 if "status" in update_data and emit_status_change:
421 await emit_status_change(
422 event_id=update_data["last_status_event_id"], # type: ignore
423 occurred=update_data["last_transitioned_status_at"],
424 pre_update_work_pool=current_work_pool,
425 work_pool=wp,
426 )
428 return updated
431@db_injector
432async def delete_work_pool(
433 db: PrefectDBInterface, session: AsyncSession, work_pool_id: UUID
434) -> bool:
435 """
436 Delete a WorkPool by id.
438 Args:
439 session (AsyncSession): A database session
440 work_pool_id (UUID): a work pool id
442 Returns:
443 bool: whether or not the WorkPool was deleted
444 """
446 work_pool = await session.get(db.WorkPool, work_pool_id)
447 if work_pool is None:
448 return False
450 queues = await read_work_queues(session=session, work_pool_id=work_pool_id)
451 async with clients.PrefectServerEventsClient() as events_client:
452 for queue in queues:
453 await events_client.emit(
454 await work_queue_deleted_event(
455 session=session, work_queue=queue, occurred=now("UTC")
456 )
457 )
459 await emit_work_pool_deleted_event(work_pool)
461 await session.execute(delete(db.WorkPool).where(db.WorkPool.id == work_pool_id))
462 return True
465@db_injector
466async def get_scheduled_flow_runs(
467 db: PrefectDBInterface,
468 session: AsyncSession,
469 work_pool_ids: Optional[List[UUID]] = None,
470 work_queue_ids: Optional[List[UUID]] = None,
471 scheduled_before: Optional[datetime.datetime] = None,
472 scheduled_after: Optional[datetime.datetime] = None,
473 limit: Optional[int] = None,
474 respect_queue_priorities: Optional[bool] = None,
475) -> Sequence[schemas.responses.WorkerFlowRunResponse]:
476 """
477 Get runs from queues in a specific work pool.
479 Args:
480 session (AsyncSession): a database session
481 work_pool_ids (List[UUID]): a list of work pool ids
482 work_queue_ids (List[UUID]): a list of work pool queue ids
483 scheduled_before (datetime.datetime): a datetime to filter runs scheduled before
484 scheduled_after (datetime.datetime): a datetime to filter runs scheduled after
485 respect_queue_priorities (bool): whether or not to respect queue priorities
486 limit (int): the maximum number of runs to return
487 db: a database interface
489 Returns:
490 List[WorkerFlowRunResponse]: the runs, as well as related work pool details
492 """
494 if respect_queue_priorities is None: 494 ↛ 497line 494 didn't jump to line 497 because the condition on line 494 was always true
495 respect_queue_priorities = True
497 return await db.queries.get_scheduled_flow_runs_from_work_pool(
498 session=session,
499 work_pool_ids=work_pool_ids,
500 work_queue_ids=work_queue_ids,
501 scheduled_before=scheduled_before,
502 scheduled_after=scheduled_after,
503 respect_queue_priorities=respect_queue_priorities,
504 limit=limit,
505 )
508# -----------------------------------------------------
509# --
510# --
511# -- Work Pool Queues
512# --
513# --
514# -----------------------------------------------------
517@db_injector
518async def create_work_queue(
519 db: PrefectDBInterface,
520 session: AsyncSession,
521 work_pool_id: UUID,
522 work_queue: schemas.actions.WorkQueueCreate,
523) -> orm_models.WorkQueue:
524 """
525 Creates a work pool queue.
527 Args:
528 session (AsyncSession): a database session
529 work_pool_id (UUID): a work pool id
530 work_queue (schemas.actions.WorkQueueCreate): a WorkQueue action model
532 Returns:
533 orm_models.WorkQueue: the newly-created WorkQueue
535 """
536 data = work_queue.model_dump(exclude={"work_pool_id"})
537 if work_queue.priority is None:
538 # Set the priority to be the first priority value that isn't already taken
539 priorities_query = sa.select(db.WorkQueue.priority).where(
540 db.WorkQueue.work_pool_id == work_pool_id
541 )
542 priorities = (await session.execute(priorities_query)).scalars().all()
544 priority = None
545 for i, p in enumerate(sorted(priorities)):
546 # if a rank was skipped (e.g. the set priority is different than the
547 # enumerated priority) then we can "take" that spot for this work
548 # queue
549 if i + 1 != p:
550 priority = i + 1
551 break
553 # otherwise take the maximum priority plus one
554 if priority is None:
555 priority = max(priorities, default=0) + 1
557 data["priority"] = priority
559 model = db.WorkQueue(**data, work_pool_id=work_pool_id)
561 session.add(model)
562 await session.flush()
563 await session.refresh(model)
565 if work_queue.priority:
566 await bulk_update_work_queue_priorities(
567 session=session,
568 work_pool_id=work_pool_id,
569 new_priorities={model.id: work_queue.priority},
570 )
572 async with clients.PrefectServerEventsClient() as events_client:
573 await events_client.emit(
574 await work_queue_created_event(
575 session=session, work_queue=model, occurred=now("UTC")
576 )
577 )
579 return model
582@db_injector
583async def bulk_update_work_queue_priorities(
584 db: PrefectDBInterface,
585 session: AsyncSession,
586 work_pool_id: UUID,
587 new_priorities: Dict[UUID, int],
588) -> None:
589 """
590 This is a brute force update of all work pool queue priorities for a given work
591 pool.
593 It loads all queues fully into memory, sorts them, and flushes the update to
594 the orm_models. The algorithm ensures that priorities are unique integers > 0, and
595 makes the minimum number of changes required to satisfy the provided
596 `new_priorities`. For example, if no queues currently have the provided
597 `new_priorities`, then they are assigned without affecting other queues. If
598 they are held by other queues, then those queues' priorities are
599 incremented as necessary.
601 Updating queue priorities is not a common operation (happens on the same scale as
602 queue modification, which is significantly less than reading from queues),
603 so while this implementation is slow, it may suffice and make up for that
604 with extreme simplicity.
605 """
607 if len(set(new_priorities.values())) != len(new_priorities): 607 ↛ 608line 607 didn't jump to line 608 because the condition on line 607 was never true
608 raise ValueError("Duplicate target priorities provided")
610 # get all the work queues, sorted by priority
611 work_queues_query = (
612 sa.select(db.WorkQueue)
613 .where(db.WorkQueue.work_pool_id == work_pool_id)
614 .order_by(db.WorkQueue.priority.asc())
615 )
616 result = await session.execute(work_queues_query)
617 all_work_queues = result.scalars().all()
619 # split the queues into those that need to be updated and those that don't
620 work_queues = [wq for wq in all_work_queues if wq.id not in new_priorities]
621 updated_queues = [wq for wq in all_work_queues if wq.id in new_priorities]
623 # update queue priorities and insert them into the appropriate place in the
624 # full list of queues
625 for queue in sorted(updated_queues, key=lambda wq: new_priorities[wq.id]): 625 ↛ anywhereline 625 didn't jump anywhere: it always raised an exception.
626 queue.priority = new_priorities[queue.id]
627 for i, wq in enumerate(work_queues):
628 if wq.priority >= new_priorities[queue.id]:
629 work_queues.insert(i, queue)
630 break
632 # walk through the queues and update their priorities such that the
633 # priorities are sequential. Do this by tracking that last priority seen and
634 # ensuring that each successive queue's priority is higher than it. This
635 # will maintain queue order and ensure increasing priorities with minimal
636 # changes.
637 last_priority = 0
638 for queue in work_queues:
639 if queue.priority <= last_priority:
640 last_priority += 1
641 queue.priority = last_priority
642 else:
643 last_priority = queue.priority
645 await session.flush()
648@db_injector
649async def read_work_queues(
650 db: PrefectDBInterface,
651 session: AsyncSession,
652 work_pool_id: UUID,
653 work_queue_filter: Optional[schemas.filters.WorkQueueFilter] = None,
654 offset: Optional[int] = None,
655 limit: Optional[int] = None,
656) -> Sequence[orm_models.WorkQueue]:
657 """
658 Read all work pool queues for a work pool. Results are ordered by ascending priority.
660 Args:
661 session (AsyncSession): a database session
662 work_pool_id (UUID): a work pool id
663 work_queue_filter: Filter criteria for work pool queues
664 offset: Query offset
665 limit: Query limit
668 Returns:
669 List[orm_models.WorkQueue]: the WorkQueues
671 """
672 query = (
673 sa.select(db.WorkQueue)
674 .where(db.WorkQueue.work_pool_id == work_pool_id)
675 .order_by(db.WorkQueue.priority.asc())
676 )
678 if work_queue_filter is not None:
679 query = query.where(work_queue_filter.as_sql_filter())
680 if offset is not None:
681 query = query.offset(offset)
682 if limit is not None:
683 query = query.limit(limit)
685 result = await session.execute(query)
686 return result.scalars().unique().all()
689@db_injector
690async def count_work_queues(
691 db: PrefectDBInterface,
692 session: AsyncSession,
693 work_pool_id: UUID,
694 work_queue_filter: Optional[schemas.filters.WorkQueueFilter] = None,
695) -> int:
696 """Count work pool queues for a work pool."""
697 query = (
698 sa.select(sa.func.count())
699 .select_from(db.WorkQueue)
700 .where(db.WorkQueue.work_pool_id == work_pool_id)
701 )
702 if work_queue_filter is not None: 702 ↛ 703line 702 didn't jump to line 703 because the condition on line 702 was never true
703 query = query.where(work_queue_filter.as_sql_filter())
704 result = await session.execute(query)
705 return result.scalar_one()
708@db_injector
709async def read_work_queue(
710 db: PrefectDBInterface,
711 session: AsyncSession,
712 work_queue_id: Union[UUID, PrefectUUID],
713) -> Optional[orm_models.WorkQueue]:
714 """
715 Read a specific work pool queue.
717 Args:
718 session (AsyncSession): a database session
719 work_queue_id (UUID): a work pool queue id
721 Returns:
722 orm_models.WorkQueue: the WorkQueue
724 """
725 return await session.get(db.WorkQueue, work_queue_id)
728@db_injector
729async def read_work_queue_by_name(
730 db: PrefectDBInterface,
731 session: AsyncSession,
732 work_pool_name: str,
733 work_queue_name: str,
734) -> Optional[orm_models.WorkQueue]:
735 """
736 Reads a WorkQueue by name.
738 Args:
739 session (AsyncSession): A database session
740 work_pool_name (str): a WorkPool name
741 work_queue_name (str): a WorkQueue name
743 Returns:
744 orm_models.WorkQueue: the WorkQueue
745 """
746 query = (
747 sa.select(db.WorkQueue)
748 .join(
749 db.WorkPool,
750 db.WorkPool.id == db.WorkQueue.work_pool_id,
751 )
752 .where(
753 db.WorkPool.name == work_pool_name,
754 db.WorkQueue.name == work_queue_name,
755 )
756 .limit(1)
757 )
758 result = await session.execute(query)
759 return result.scalar()
762@db_injector
763async def update_work_queue(
764 db: PrefectDBInterface,
765 session: AsyncSession,
766 work_queue_id: UUID,
767 work_queue: schemas.actions.WorkQueueUpdate,
768 emit_status_change: Optional[
769 Callable[[orm_models.WorkQueue], Awaitable[None]]
770 ] = None,
771 default_status: WorkQueueStatus = WorkQueueStatus.NOT_READY,
772) -> bool:
773 """
774 Update a work pool queue.
776 Args:
777 session (AsyncSession): a database session
778 work_queue_id (UUID): a work pool queue ID
779 work_queue (schemas.actions.WorkQueueUpdate): a WorkQueue model
780 emit_status_change: function to call when work queue
781 status is changed
783 Returns:
784 bool: whether or not the WorkQueue was updated
786 """
787 from prefect.server.models.work_queues import is_last_polled_recent
789 update_values = work_queue.model_dump_for_orm(exclude_unset=True)
791 current_work_queue = await session.get(db.WorkQueue, work_queue_id)
792 if current_work_queue is None:
793 return False
794 session.expunge(current_work_queue)
796 if "is_paused" in update_values:
797 # Only update the status to paused if it's not already paused. This ensures a work queue that is already
798 # paused will not get a status update if it's paused again
799 if (
800 update_values.get("is_paused")
801 and current_work_queue.status != WorkQueueStatus.PAUSED
802 ):
803 update_values["status"] = WorkQueueStatus.PAUSED
805 # If unpausing, only update status if it's currently paused. This ensures a work queue that is already
806 # unpaused will not get a status update if it's unpaused again
807 if (
808 update_values.get("is_paused") is False
809 and current_work_queue.status == WorkQueueStatus.PAUSED
810 ):
811 # Default status if unpaused
812 update_values["status"] = default_status
814 # Determine source of last_polled: update_data or database
815 if "last_polled" in update_values:
816 last_polled = update_values["last_polled"]
817 else:
818 last_polled = current_work_queue.last_polled
820 # Check if last polled is recent and set status to READY if so
821 if is_last_polled_recent(last_polled):
822 update_values["status"] = schemas.statuses.WorkQueueStatus.READY
824 update_stmt = (
825 sa.update(db.WorkQueue)
826 .where(db.WorkQueue.id == work_queue_id)
827 .values(update_values)
828 )
829 result = await session.execute(update_stmt)
831 updated = result.rowcount > 0
833 if updated:
834 updated_work_queue = await session.get(db.WorkQueue, work_queue_id)
835 assert updated_work_queue is not None
836 assert current_work_queue is not updated_work_queue
838 # Fields that should trigger update events (user-updatable fields)
839 WORK_QUEUE_EVENT_FIELDS = {
840 "name",
841 "description",
842 "concurrency_limit",
843 "priority",
844 # Exclude "is_paused" - handled with status
845 # Exclude "last_polled" - usually auto-updated
846 # Exclude "filter" - deprecated
847 }
849 changed_fields = {}
850 for field in update_values.keys():
851 if field not in WORK_QUEUE_EVENT_FIELDS or field == "status":
852 continue
854 old_value = getattr(current_work_queue, field, None)
855 new_value = getattr(updated_work_queue, field, None)
857 if old_value != new_value:
858 changed_fields[field] = {
859 "from": old_value,
860 "to": new_value,
861 }
863 if changed_fields:
864 from prefect.server.models.work_queues import emit_work_queue_updated_event
866 await emit_work_queue_updated_event(
867 session=session,
868 work_queue=updated_work_queue,
869 changed_fields=changed_fields,
870 )
872 if "priority" in update_values or "status" in update_values:
873 if "priority" in update_values:
874 await bulk_update_work_queue_priorities(
875 session,
876 work_pool_id=updated_work_queue.work_pool_id,
877 new_priorities={work_queue_id: update_values["priority"]},
878 )
880 if "status" in update_values and emit_status_change:
881 await emit_status_change(updated_work_queue)
883 return updated
886@db_injector
887async def delete_work_queue(
888 db: PrefectDBInterface,
889 session: AsyncSession,
890 work_queue_id: UUID,
891) -> bool:
892 """
893 Delete a work pool queue.
895 Args:
896 session (AsyncSession): a database session
897 work_queue_id (UUID): a work pool queue ID
899 Returns:
900 bool: whether or not the WorkQueue was deleted
902 """
903 work_queue = await session.get(db.WorkQueue, work_queue_id)
904 if work_queue is None:
905 return False
907 async with clients.PrefectServerEventsClient() as events_client:
908 await events_client.emit(
909 await work_queue_deleted_event(
910 session=session, work_queue=work_queue, occurred=now("UTC")
911 )
912 )
914 await session.delete(work_queue)
915 try:
916 await session.flush()
918 # if an error was raised, check if the user tried to delete a default queue
919 except sa.exc.IntegrityError as exc:
920 if "foreign key constraint" in str(exc).lower():
921 raise ValueError("Can't delete a pool's default queue.")
922 raise
924 await bulk_update_work_queue_priorities(
925 session,
926 work_pool_id=work_queue.work_pool_id,
927 new_priorities={},
928 )
929 return True
932# -----------------------------------------------------
933# --
934# --
935# -- Workers
936# --
937# --
938# -----------------------------------------------------
941@db_injector
942async def read_workers(
943 db: PrefectDBInterface,
944 session: AsyncSession,
945 work_pool_id: UUID,
946 worker_filter: Optional[schemas.filters.WorkerFilter] = None,
947 limit: Optional[int] = None,
948 offset: Optional[int] = None,
949) -> Sequence[orm_models.Worker]:
950 query = (
951 sa.select(db.Worker)
952 .where(db.Worker.work_pool_id == work_pool_id)
953 .order_by(db.Worker.last_heartbeat_time.desc())
954 .limit(limit)
955 )
957 if worker_filter:
958 query = query.where(worker_filter.as_sql_filter())
960 if limit is not None: 960 ↛ 963line 960 didn't jump to line 963 because the condition on line 960 was always true
961 query = query.limit(limit)
963 if offset is not None: 963 ↛ 966line 963 didn't jump to line 966 because the condition on line 963 was always true
964 query = query.offset(offset)
966 result = await session.execute(query)
967 return result.scalars().all()
970@db_injector
971async def read_worker_by_name(
972 db: PrefectDBInterface,
973 session: AsyncSession,
974 work_pool_id: UUID,
975 worker_name: str,
976) -> Optional[orm_models.Worker]:
977 query = (
978 sa.select(db.Worker)
979 .where(
980 db.Worker.work_pool_id == work_pool_id,
981 db.Worker.name == worker_name,
982 )
983 .limit(1)
984 )
985 result = await session.execute(query)
986 return result.scalar()
989@db_injector
990async def worker_heartbeat(
991 db: PrefectDBInterface,
992 session: AsyncSession,
993 work_pool_id: UUID,
994 worker_name: str,
995 heartbeat_interval_seconds: Optional[int] = None,
996) -> bool:
997 """
998 Record a worker process heartbeat.
1000 Args:
1001 session (AsyncSession): a database session
1002 work_pool_id (UUID): a work pool ID
1003 worker_name (str): a worker name
1005 Returns:
1006 bool: whether or not the worker was updated
1008 """
1009 right_now = now("UTC")
1010 # Values that won't change between heart beats
1011 base_values = dict(
1012 work_pool_id=work_pool_id,
1013 name=worker_name,
1014 )
1015 # Values that can and will change between heartbeats
1016 update_values = dict(
1017 last_heartbeat_time=right_now,
1018 status=schemas.statuses.WorkerStatus.ONLINE,
1019 )
1020 if heartbeat_interval_seconds is not None:
1021 update_values["heartbeat_interval_seconds"] = heartbeat_interval_seconds
1023 insert_stmt = (
1024 db.queries.insert(db.Worker)
1025 .values(**base_values, **update_values)
1026 .on_conflict_do_update(
1027 index_elements=[
1028 db.Worker.work_pool_id,
1029 db.Worker.name,
1030 ],
1031 set_=update_values,
1032 )
1033 )
1035 result = await session.execute(insert_stmt)
1036 return result.rowcount > 0
1039async def record_worker_heartbeat(
1040 session: AsyncSession,
1041 work_pool: orm_models.WorkPool,
1042 worker_name: str,
1043 heartbeat_interval_seconds: Optional[int] = None,
1044 emit_status_change: Optional[
1045 Callable[
1046 [UUID, DateTime, orm_models.WorkPool, orm_models.WorkPool],
1047 Awaitable[None],
1048 ]
1049 ] = None,
1050 return_worker: bool = False,
1051) -> Optional[orm_models.Worker]:
1052 await worker_heartbeat(
1053 session=session,
1054 work_pool_id=work_pool.id,
1055 worker_name=worker_name,
1056 heartbeat_interval_seconds=heartbeat_interval_seconds,
1057 )
1059 if work_pool.status == schemas.statuses.WorkPoolStatus.NOT_READY:
1060 await update_work_pool(
1061 session=session,
1062 work_pool_id=work_pool.id,
1063 work_pool=schemas.internal.InternalWorkPoolUpdate(
1064 status=schemas.statuses.WorkPoolStatus.READY
1065 ),
1066 emit_status_change=emit_status_change,
1067 )
1069 if not return_worker:
1070 return None
1072 worker = await read_worker_by_name(
1073 session=session,
1074 work_pool_id=work_pool.id,
1075 worker_name=worker_name,
1076 )
1077 assert worker is not None
1078 return worker
1081@db_injector
1082async def delete_worker(
1083 db: PrefectDBInterface,
1084 session: AsyncSession,
1085 work_pool_id: UUID,
1086 worker_name: str,
1087) -> bool:
1088 """
1089 Delete a work pool's worker.
1091 Args:
1092 session (AsyncSession): a database session
1093 work_pool_id (UUID): a work pool ID
1094 worker_name (str): a worker name
1096 Returns:
1097 bool: whether or not the Worker was deleted
1099 """
1100 result = await session.execute(
1101 delete(db.Worker).where(
1102 db.Worker.work_pool_id == work_pool_id,
1103 db.Worker.name == worker_name,
1104 )
1105 )
1107 return result.rowcount > 0
1110# Work-pool scheduler (get-runs-from-worker-queues.sql.jinja) counts only
1111# PENDING and RUNNING against pool/queue concurrency limits.
1112WORK_POOL_SLOT_OCCUPYING_STATES = [
1113 schemas.states.StateType.PENDING,
1114 schemas.states.StateType.RUNNING,
1115]
1117# Work-queue scheduler (query_components.py) also counts CANCELLING.
1118WORK_QUEUE_SLOT_OCCUPYING_STATES = [
1119 schemas.states.StateType.PENDING,
1120 schemas.states.StateType.RUNNING,
1121 schemas.states.StateType.CANCELLING,
1122]
1124# Union of both for the slot_acquired_at subquery, which needs to find
1125# the earliest entry into any slot-occupying state regardless of context.
1126ALL_SLOT_OCCUPYING_STATES = [
1127 schemas.states.StateType.PENDING,
1128 schemas.states.StateType.RUNNING,
1129 schemas.states.StateType.CANCELLING,
1130]
1133def _slot_acquired_at_subquery(
1134 db: PrefectDBInterface,
1135) -> sa.ScalarSelect:
1136 """Correlated subquery returning when the current slot-occupying sequence began.
1138 For a run that has been retried or rescheduled, this finds the start of the
1139 *current* slot-occupying sequence — not the first-ever one. It does this by
1140 finding the latest non-slot-occupying state and then taking the earliest
1141 slot-occupying state after that point.
1142 """
1143 # Latest non-slot-occupying state timestamp (e.g. SCHEDULED, FAILED before retry)
1144 last_non_slot_state = (
1145 select(sa.func.max(db.FlowRunState.timestamp))
1146 .where(
1147 db.FlowRunState.flow_run_id == db.FlowRun.id,
1148 db.FlowRunState.type.notin_(ALL_SLOT_OCCUPYING_STATES),
1149 )
1150 .correlate(db.FlowRun)
1151 .scalar_subquery()
1152 )
1154 # Earliest slot-occupying state at or after the last non-slot state.
1155 # Uses >= so that a slot-occupying state with the same timestamp as the
1156 # preceding non-slot state (possible with imported/manual timestamps or
1157 # coarse precision) is still recognized as the current attempt.
1158 # If there was never a non-slot state, returns the earliest slot-occupying
1159 # state overall (correct for runs that started directly in a slot-occupying state).
1160 return (
1161 select(sa.func.min(db.FlowRunState.timestamp))
1162 .where(
1163 db.FlowRunState.flow_run_id == db.FlowRun.id,
1164 db.FlowRunState.type.in_(ALL_SLOT_OCCUPYING_STATES),
1165 sa.or_(
1166 last_non_slot_state.is_(None),
1167 db.FlowRunState.timestamp >= last_non_slot_state,
1168 ),
1169 )
1170 .correlate(db.FlowRun)
1171 .scalar_subquery()
1172 .label("slot_acquired_at")
1173 )
1176def _work_pool_slot_holder_filter(
1177 db: PrefectDBInterface, work_pool_id: UUID
1178) -> sa.ColumnElement:
1179 """Common WHERE clause for work-pool slot holder queries."""
1180 return sa.and_(
1181 db.WorkQueue.work_pool_id == work_pool_id,
1182 db.FlowRun.state_type.in_(WORK_POOL_SLOT_OCCUPYING_STATES),
1183 )
1186@db_injector
1187async def count_work_pool_slot_holders(
1188 db: PrefectDBInterface,
1189 session: AsyncSession,
1190 work_pool_id: UUID,
1191) -> int:
1192 """Counts flow runs in slot-occupying states for a work pool."""
1193 query = (
1194 select(sa.func.count())
1195 .select_from(db.FlowRun)
1196 .join(db.WorkQueue, db.FlowRun.work_queue_id == db.WorkQueue.id)
1197 .where(_work_pool_slot_holder_filter(db, work_pool_id))
1198 )
1199 result = await session.execute(query)
1200 return result.scalar_one()
1203@db_injector
1204async def get_work_pool_slot_holders(
1205 db: PrefectDBInterface,
1206 session: AsyncSession,
1207 work_pool_id: UUID,
1208 work_queue_ids: Optional[List[UUID]] = None,
1209 flow_run_limit: Optional[int] = None,
1210) -> Sequence[tuple[orm_models.FlowRun, Optional[DateTime]]]:
1211 """Returns flow runs in slot-occupying states for a work pool.
1213 Each result is a tuple of (FlowRun, slot_acquired_at) where
1214 slot_acquired_at is when the current slot-occupying sequence began.
1216 Args:
1217 work_pool_id: The work pool to query.
1218 work_queue_ids: If provided, only return runs for these queues.
1219 flow_run_limit: If provided, cap results per work_queue_id.
1220 """
1221 slot_acquired_at = _slot_acquired_at_subquery(db)
1222 # Matches the work-pool scheduler (get-runs-from-worker-queues.sql.jinja):
1223 # - Joins on work_queue_id (not name) — fr.work_queue_id = wq.id
1224 # - Counts only PENDING and RUNNING (not CANCELLING)
1225 # Does NOT filter on queue pause status: the Postgres and SQLite scheduler
1226 # templates disagree on this (Postgres excludes paused queues, SQLite
1227 # doesn't), and as a status/debugging endpoint we want to show all runs
1228 # that are actually consuming resources regardless of pause state.
1229 query = (
1230 select(db.FlowRun, slot_acquired_at)
1231 .join(db.WorkQueue, db.FlowRun.work_queue_id == db.WorkQueue.id)
1232 .where(_work_pool_slot_holder_filter(db, work_pool_id))
1233 .order_by(db.FlowRun.id)
1234 )
1235 if work_queue_ids is not None: 1235 ↛ 1237line 1235 didn't jump to line 1237 because the condition on line 1235 was always true
1236 query = query.where(db.FlowRun.work_queue_id.in_(work_queue_ids))
1237 result = await session.execute(query)
1238 rows = result.all()
1240 if flow_run_limit is not None and work_queue_ids is not None:
1241 # Apply per-queue flow_run_limit in Python (simpler than SQL windowing)
1242 from collections import Counter
1244 counts: Counter[UUID] = Counter()
1245 limited: list[tuple] = []
1246 for run, sa_val in rows:
1247 qid = run.work_queue_id
1248 if counts[qid] < flow_run_limit:
1249 limited.append((run, sa_val))
1250 counts[qid] += 1
1251 return limited
1253 return rows
1256@db_injector
1257async def count_work_pool_slot_holders_by_queue(
1258 db: PrefectDBInterface,
1259 session: AsyncSession,
1260 work_pool_id: UUID,
1261) -> dict[UUID, int]:
1262 """Returns `{work_queue_id: count}` for slot-holding runs in a pool."""
1263 query = (
1264 select(db.FlowRun.work_queue_id, sa.func.count())
1265 .join(db.WorkQueue, db.FlowRun.work_queue_id == db.WorkQueue.id)
1266 .where(_work_pool_slot_holder_filter(db, work_pool_id))
1267 .group_by(db.FlowRun.work_queue_id)
1268 )
1269 result = await session.execute(query)
1270 return dict(result.all())
1273def _work_queue_slot_holder_filter(
1274 db: PrefectDBInterface, work_queue_id: UUID, queue_name_subquery: sa.ScalarSelect
1275) -> sa.ColumnElement:
1276 """Common WHERE clause for work-queue slot holder queries."""
1277 return sa.and_(
1278 sa.or_(
1279 db.FlowRun.work_queue_id == work_queue_id,
1280 sa.and_(
1281 db.FlowRun.work_queue_id.is_(None),
1282 db.FlowRun.work_queue_name == queue_name_subquery,
1283 ),
1284 ),
1285 db.FlowRun.state_type.in_(WORK_QUEUE_SLOT_OCCUPYING_STATES),
1286 )
1289@db_injector
1290async def count_work_queue_slot_holders(
1291 db: PrefectDBInterface,
1292 session: AsyncSession,
1293 work_queue_id: UUID,
1294) -> int:
1295 """Counts flow runs in slot-occupying states for a single work queue."""
1296 queue_name_subquery = (
1297 select(db.WorkQueue.name)
1298 .where(db.WorkQueue.id == work_queue_id)
1299 .scalar_subquery()
1300 )
1301 query = (
1302 select(sa.func.count())
1303 .select_from(db.FlowRun)
1304 .where(_work_queue_slot_holder_filter(db, work_queue_id, queue_name_subquery))
1305 )
1306 result = await session.execute(query)
1307 return result.scalar_one()
1310@db_injector
1311async def get_work_queue_slot_holders(
1312 db: PrefectDBInterface,
1313 session: AsyncSession,
1314 work_queue_id: UUID,
1315 offset: Optional[int] = None,
1316 limit: Optional[int] = None,
1317) -> Sequence[tuple[orm_models.FlowRun, Optional[DateTime]]]:
1318 """Returns flow runs in slot-occupying states for a single work queue.
1320 Each result is a tuple of (FlowRun, slot_acquired_at) where
1321 slot_acquired_at is when the current slot-occupying sequence began.
1322 """
1323 slot_acquired_at = _slot_acquired_at_subquery(db)
1324 # Matches the work-queue scheduler (query_components.py):
1325 # - Joins on work_queue_name (FlowRun.work_queue_name == WorkQueue.name)
1326 # - Counts PENDING, RUNNING, and CANCELLING
1327 queue_name_subquery = (
1328 select(db.WorkQueue.name)
1329 .where(db.WorkQueue.id == work_queue_id)
1330 .scalar_subquery()
1331 )
1332 query = (
1333 select(db.FlowRun, slot_acquired_at)
1334 .where(_work_queue_slot_holder_filter(db, work_queue_id, queue_name_subquery))
1335 .order_by(db.FlowRun.id)
1336 )
1337 if offset is not None: 1337 ↛ 1339line 1337 didn't jump to line 1339 because the condition on line 1337 was always true
1338 query = query.offset(offset)
1339 if limit is not None: 1339 ↛ 1341line 1339 didn't jump to line 1341 because the condition on line 1339 was always true
1340 query = query.limit(limit)
1341 result = await session.execute(query)
1342 return result.all()
1345async def emit_work_pool_updated_event(
1346 session: AsyncSession,
1347 work_pool: orm_models.WorkPool,
1348 changed_fields: Dict[str, Dict[str, Any]],
1349) -> None:
1350 """Emit an event when work pool fields are updated."""
1351 if not changed_fields:
1352 return
1354 async with PrefectServerEventsClient() as events_client:
1355 await events_client.emit(
1356 await work_pool_updated_event(
1357 session=session,
1358 work_pool=work_pool,
1359 changed_fields=changed_fields,
1360 occurred=now("UTC"),
1361 )
1362 )
1365async def emit_work_pool_created_event(work_pool: orm_models.WorkPool) -> None:
1366 """Emit an event when a work pool is created."""
1367 async with clients.PrefectServerEventsClient() as events_client:
1368 await events_client.emit(
1369 await work_pool_created_event(work_pool=work_pool, occurred=now("UTC"))
1370 )
1373async def emit_work_pool_deleted_event(work_pool: orm_models.WorkPool) -> None:
1374 """Emit an event when a work pool is deleted."""
1375 async with clients.PrefectServerEventsClient() as events_client:
1376 await events_client.emit(
1377 await work_pool_deleted_event(work_pool=work_pool, occurred=now("UTC"))
1378 )
1381async def emit_work_pool_status_event(
1382 event_id: UUID,
1383 occurred: DateTime,
1384 pre_update_work_pool: Optional[orm_models.WorkPool],
1385 work_pool: orm_models.WorkPool,
1386) -> None:
1387 if not work_pool.status: 1387 ↛ 1388line 1387 didn't jump to line 1388 because the condition on line 1387 was never true
1388 return
1390 async with PrefectServerEventsClient() as events_client:
1391 await events_client.emit(
1392 await work_pool_status_event(
1393 event_id=event_id,
1394 occurred=occurred,
1395 pre_update_work_pool=pre_update_work_pool,
1396 work_pool=work_pool,
1397 )
1398 )