Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/models/task_runs.py: 76%
161 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 task run ORM objects.
3Intended for internal use by the Prefect REST API.
4"""
6import contextlib
7from typing import (
8 TYPE_CHECKING,
9 Any,
10 Dict,
11 Optional,
12 Sequence,
13 Type,
14 TypeVar,
15 Union,
16 cast,
17)
18from uuid import UUID
20import sqlalchemy as sa
21from sqlalchemy import delete, select
22from sqlalchemy.ext.asyncio import AsyncSession
23from sqlalchemy.sql import Select
25import prefect.server.models as models
26import prefect.server.schemas as schemas
27from prefect.logging import get_logger
28from prefect.server.database import PrefectDBInterface, db_injector, orm_models
29from prefect.server.exceptions import ObjectNotFoundError
30from prefect.server.orchestration.core_policy import (
31 BackgroundTaskPolicy,
32 MinimalTaskPolicy,
33)
34from prefect.server.orchestration.global_policy import GlobalTaskPolicy
35from prefect.server.orchestration.policies import (
36 TaskRunOrchestrationPolicy,
37)
38from prefect.server.orchestration.rules import TaskOrchestrationContext
39from prefect.server.schemas.responses import OrchestrationResult
40from prefect.types._datetime import now
42if TYPE_CHECKING: 42 ↛ 43line 42 didn't jump to line 43 because the condition on line 42 was never true
43 import logging
45T = TypeVar("T", bound=tuple[Any, ...])
47logger: "logging.Logger" = get_logger("server")
50@db_injector
51async def create_task_run(
52 db: PrefectDBInterface,
53 session: AsyncSession,
54 task_run: schemas.core.TaskRun,
55 orchestration_parameters: Optional[Dict[str, Any]] = None,
56) -> orm_models.TaskRun:
57 """
58 Creates a new task run.
60 If a task run with the same flow_run_id, task_key, and dynamic_key already exists,
61 the existing task run will be returned. If the provided task run has a state
62 attached, it will also be created.
64 Args:
65 session: a database session
66 task_run: a task run model
68 Returns:
69 orm_models.TaskRun: the newly-created or existing task run
70 """
72 right_now = now("UTC")
73 model: Union[orm_models.TaskRun, None]
75 task_run.labels = await with_system_labels_for_task_run(
76 session=session, task_run=task_run
77 )
79 # if a dynamic key exists, we need to guard against conflicts
80 if task_run.flow_run_id: 80 ↛ 81line 80 didn't jump to line 81 because the condition on line 80 was never true
81 insert_stmt = (
82 db.queries.insert(db.TaskRun)
83 .values(
84 created=right_now,
85 **task_run.model_dump_for_orm(
86 exclude={"state", "created"}, exclude_unset=True
87 ),
88 )
89 .on_conflict_do_nothing(
90 index_elements=db.orm.task_run_unique_upsert_columns,
91 )
92 )
93 await session.execute(insert_stmt)
95 query = (
96 sa.select(db.TaskRun)
97 .where(
98 sa.and_(
99 db.TaskRun.flow_run_id == task_run.flow_run_id,
100 db.TaskRun.task_key == task_run.task_key,
101 db.TaskRun.dynamic_key == task_run.dynamic_key,
102 )
103 )
104 .limit(1)
105 .execution_options(populate_existing=True)
106 )
107 result = await session.execute(query)
108 model = result.scalar_one()
109 else:
110 # Upsert on (task_key, dynamic_key) application logic.
111 query = (
112 sa.select(db.TaskRun)
113 .where(
114 sa.and_(
115 db.TaskRun.flow_run_id.is_(None),
116 db.TaskRun.task_key == task_run.task_key,
117 db.TaskRun.dynamic_key == task_run.dynamic_key,
118 )
119 )
120 .limit(1)
121 .execution_options(populate_existing=True)
122 )
124 result = await session.execute(query)
125 model = result.scalar()
127 if model is None:
128 model = db.TaskRun(
129 created=right_now,
130 **task_run.model_dump_for_orm(
131 exclude={"state", "created"}, exclude_unset=True
132 ),
133 state=None,
134 )
135 session.add(model)
136 await session.flush()
138 if model.created == right_now and task_run.state:
139 await models.task_runs.set_task_run_state(
140 session=session,
141 task_run_id=model.id,
142 state=task_run.state,
143 force=True,
144 orchestration_parameters=orchestration_parameters,
145 )
147 return model
150@db_injector
151async def update_task_run(
152 db: PrefectDBInterface,
153 session: AsyncSession,
154 task_run_id: UUID,
155 task_run: schemas.actions.TaskRunUpdate,
156) -> bool:
157 """
158 Updates a task run.
160 Args:
161 session: a database session
162 task_run_id: the task run id to update
163 task_run: a task run model
165 Returns:
166 bool: whether or not matching rows were found to update
167 """
168 update_stmt = (
169 sa.update(db.TaskRun)
170 .where(db.TaskRun.id == task_run_id)
171 # exclude_unset=True allows us to only update values provided by
172 # the user, ignoring any defaults on the model
173 .values(**task_run.model_dump_for_orm(exclude_unset=True))
174 )
175 result = await session.execute(update_stmt)
176 return result.rowcount > 0
179@db_injector
180async def read_task_run(
181 db: PrefectDBInterface, session: AsyncSession, task_run_id: UUID
182) -> Union[orm_models.TaskRun, None]:
183 """
184 Read a task run by id.
186 Args:
187 session: a database session
188 task_run_id: the task run id
190 Returns:
191 orm_models.TaskRun: the task run
192 """
194 model = await session.get(db.TaskRun, task_run_id)
195 return model
198@db_injector
199async def read_task_run_with_flow_run_name(
200 db: PrefectDBInterface, session: AsyncSession, task_run_id: UUID
201) -> Union[orm_models.TaskRun, None]:
202 """
203 Read a task run by id.
205 Args:
206 session: a database session
207 task_run_id: the task run id
209 Returns:
210 orm_models.TaskRun: the task run with the flow run name
211 """
213 result = await session.execute(
214 select(orm_models.TaskRun, orm_models.FlowRun.name.label("flow_run_name"))
215 .outerjoin(
216 orm_models.FlowRun, orm_models.TaskRun.flow_run_id == orm_models.FlowRun.id
217 )
218 .where(orm_models.TaskRun.id == task_run_id)
219 )
220 row = result.first()
221 if not row:
222 return None
224 task_run = row[0]
225 flow_run_name = row[1]
226 if flow_run_name:
227 setattr(task_run, "flow_run_name", flow_run_name)
228 return task_run
231async def _apply_task_run_filters(
232 db: PrefectDBInterface,
233 query: Select[T],
234 flow_filter: Optional[schemas.filters.FlowFilter] = None,
235 flow_run_filter: Optional[schemas.filters.FlowRunFilter] = None,
236 task_run_filter: Optional[schemas.filters.TaskRunFilter] = None,
237 deployment_filter: Optional[schemas.filters.DeploymentFilter] = None,
238 work_pool_filter: Optional[schemas.filters.WorkPoolFilter] = None,
239 work_queue_filter: Optional[schemas.filters.WorkQueueFilter] = None,
240) -> Select[T]:
241 """
242 Applies filters to a task run query as a combination of EXISTS subqueries.
243 """
245 if task_run_filter:
246 query = query.where(task_run_filter.as_sql_filter())
248 # Return a simplified query in the case that the request is ONLY asking to filter on flow_run_id (and task_run_filter)
249 # In this case there's no need to generate the complex EXISTS subqueries; the generated query here is much more efficient
250 if (
251 flow_run_filter
252 and flow_run_filter.only_filters_on_id()
253 and flow_run_filter.id
254 and flow_run_filter.id.any_
255 and not any(
256 [flow_filter, deployment_filter, work_pool_filter, work_queue_filter]
257 )
258 ):
259 query = query.where(db.TaskRun.flow_run_id.in_(flow_run_filter.id.any_))
261 return query
263 if (
264 flow_filter
265 or flow_run_filter
266 or deployment_filter
267 or work_pool_filter
268 or work_queue_filter
269 ):
270 exists_clause = select(db.FlowRun).where(
271 db.FlowRun.id == db.TaskRun.flow_run_id
272 )
274 if flow_run_filter:
275 exists_clause = exists_clause.where(flow_run_filter.as_sql_filter())
277 if flow_filter:
278 exists_clause = exists_clause.join(
279 db.Flow,
280 db.Flow.id == db.FlowRun.flow_id,
281 ).where(flow_filter.as_sql_filter())
283 if deployment_filter:
284 exists_clause = exists_clause.join(
285 db.Deployment,
286 db.Deployment.id == db.FlowRun.deployment_id,
287 ).where(deployment_filter.as_sql_filter())
289 if work_queue_filter: 289 ↛ 290line 289 didn't jump to line 290 because the condition on line 289 was never true
290 exists_clause = exists_clause.join(
291 db.WorkQueue,
292 db.WorkQueue.id == db.FlowRun.work_queue_id,
293 ).where(work_queue_filter.as_sql_filter())
295 if work_pool_filter: 295 ↛ 296line 295 didn't jump to line 296 because the condition on line 295 was never true
296 exists_clause = exists_clause.join(
297 db.WorkPool,
298 sa.and_(
299 db.WorkPool.id == db.WorkQueue.work_pool_id,
300 db.WorkQueue.id == db.FlowRun.work_queue_id,
301 ),
302 ).where(work_pool_filter.as_sql_filter())
304 query = query.where(exists_clause.exists())
306 return query
309@db_injector
310async def read_task_runs(
311 db: PrefectDBInterface,
312 session: AsyncSession,
313 flow_filter: Optional[schemas.filters.FlowFilter] = None,
314 flow_run_filter: Optional[schemas.filters.FlowRunFilter] = None,
315 task_run_filter: Optional[schemas.filters.TaskRunFilter] = None,
316 deployment_filter: Optional[schemas.filters.DeploymentFilter] = None,
317 offset: Optional[int] = None,
318 limit: Optional[int] = None,
319 sort: schemas.sorting.TaskRunSort = schemas.sorting.TaskRunSort.ID_DESC,
320) -> Sequence[orm_models.TaskRun]:
321 """
322 Read task runs.
324 Args:
325 session: a database session
326 flow_filter: only select task runs whose flows match these filters
327 flow_run_filter: only select task runs whose flow runs match these filters
328 task_run_filter: only select task runs that match these filters
329 deployment_filter: only select task runs whose deployments match these filters
330 offset: Query offset
331 limit: Query limit
332 sort: Query sort
334 Returns:
335 List[orm_models.TaskRun]: the task runs
336 """
338 query = select(db.TaskRun).order_by(*sort.as_sql_sort())
340 query = await _apply_task_run_filters(
341 db,
342 query,
343 flow_filter=flow_filter,
344 flow_run_filter=flow_run_filter,
345 task_run_filter=task_run_filter,
346 deployment_filter=deployment_filter,
347 )
349 if offset is not None:
350 query = query.offset(offset)
352 if limit is not None:
353 query = query.limit(limit)
355 logger.debug(f"In read_task_runs, query generated is:\n{query}")
356 result = await session.execute(query)
357 return result.scalars().unique().all()
360@db_injector
361async def count_task_runs(
362 db: PrefectDBInterface,
363 session: AsyncSession,
364 flow_filter: Optional[schemas.filters.FlowFilter] = None,
365 flow_run_filter: Optional[schemas.filters.FlowRunFilter] = None,
366 task_run_filter: Optional[schemas.filters.TaskRunFilter] = None,
367 deployment_filter: Optional[schemas.filters.DeploymentFilter] = None,
368) -> int:
369 """
370 Count task runs.
372 Args:
373 session: a database session
374 flow_filter: only count task runs whose flows match these filters
375 flow_run_filter: only count task runs whose flow runs match these filters
376 task_run_filter: only count task runs that match these filters
377 deployment_filter: only count task runs whose deployments match these filters
378 Returns:
379 int: count of task runs
380 """
382 if flow_filter or flow_run_filter or deployment_filter:
383 query = select(sa.func.count(None)).select_from(db.TaskRun)
384 query = query.join(db.FlowRun, db.TaskRun.flow_run_id == db.FlowRun.id)
386 if flow_run_filter:
387 query = query.where(flow_run_filter.as_sql_filter())
389 if flow_filter:
390 query = query.join(db.Flow, db.Flow.id == db.FlowRun.flow_id)
391 query = query.where(flow_filter.as_sql_filter())
393 if deployment_filter:
394 query = query.join(
395 db.Deployment, db.Deployment.id == db.FlowRun.deployment_id
396 )
397 query = query.where(deployment_filter.as_sql_filter())
399 if task_run_filter:
400 query = query.where(task_run_filter.as_sql_filter())
401 else:
402 query = select(sa.func.count(None)).select_from(db.TaskRun)
404 query = await _apply_task_run_filters(
405 db,
406 query,
407 flow_filter=flow_filter,
408 flow_run_filter=flow_run_filter,
409 task_run_filter=task_run_filter,
410 deployment_filter=deployment_filter,
411 )
413 result = await session.execute(query)
414 return result.scalar_one()
417@db_injector
418async def count_task_runs_by_state(
419 db: PrefectDBInterface,
420 session: AsyncSession,
421 flow_filter: Optional[schemas.filters.FlowFilter] = None,
422 flow_run_filter: Optional[schemas.filters.FlowRunFilter] = None,
423 task_run_filter: Optional[schemas.filters.TaskRunFilter] = None,
424 deployment_filter: Optional[schemas.filters.DeploymentFilter] = None,
425) -> schemas.states.CountByState:
426 """
427 Count task runs by state.
429 Args:
430 session: a database session
431 flow_filter: only count task runs whose flows match these filters
432 flow_run_filter: only count task runs whose flow runs match these filters
433 task_run_filter: only count task runs that match these filters
434 deployment_filter: only count task runs whose deployments match these filters
435 Returns:
436 schemas.states.CountByState: count of task runs by state
437 """
439 base_query = (
440 select(db.TaskRun.state_type, sa.func.count(None).label("count"))
441 .select_from(db.TaskRun)
442 .group_by(db.TaskRun.state_type)
443 .where(db.TaskRun.state_type.isnot(None))
444 )
446 query = await _apply_task_run_filters(
447 db,
448 base_query,
449 flow_filter=flow_filter,
450 flow_run_filter=flow_run_filter,
451 task_run_filter=task_run_filter,
452 deployment_filter=deployment_filter,
453 )
455 result = await session.execute(query)
457 counts = schemas.states.CountByState()
459 for row in result:
460 setattr(counts, row.state_type, row.count)
462 return counts
465@db_injector
466async def delete_task_run(
467 db: PrefectDBInterface, session: AsyncSession, task_run_id: UUID
468) -> bool:
469 """
470 Delete a task run by id.
472 Args:
473 session: a database session
474 task_run_id: the task run id to delete
476 Returns:
477 bool: whether or not the task run was deleted
478 """
480 result = await session.execute(
481 delete(db.TaskRun).where(db.TaskRun.id == task_run_id)
482 )
483 return result.rowcount > 0
486async def set_task_run_state(
487 session: AsyncSession,
488 task_run_id: UUID,
489 state: schemas.states.State,
490 force: bool = False,
491 task_policy: Optional[Type[TaskRunOrchestrationPolicy]] = None,
492 orchestration_parameters: Optional[Dict[str, Any]] = None,
493) -> OrchestrationResult:
494 """
495 Creates a new orchestrated task run state.
497 Setting a new state on a run is the one of the principal actions that is governed by
498 Prefect's orchestration logic. Setting a new run state will not guarantee creation,
499 but instead trigger orchestration rules to govern the proposed `state` input. If
500 the state is considered valid, it will be written to the database. Otherwise, a
501 it's possible a different state, or no state, will be created. A `force` flag is
502 supplied to bypass a subset of orchestration logic.
504 Args:
505 session: a database session
506 task_run_id: the task run id
507 state: a task run state model
508 force: if False, orchestration rules will be applied that may alter or prevent
509 the state transition. If True, orchestration rules are not applied.
511 Returns:
512 OrchestrationResult object
513 """
515 # load the task run
516 run = await models.task_runs.read_task_run(session=session, task_run_id=task_run_id)
518 if not run: 518 ↛ 519line 518 didn't jump to line 519 because the condition on line 518 was never true
519 raise ObjectNotFoundError(f"Task run with id {task_run_id} not found")
521 initial_state = run.state.as_state() if run.state else None
522 initial_state_type = initial_state.type if initial_state else None
523 proposed_state_type = state.type if state else None
524 intended_transition = (initial_state_type, proposed_state_type)
526 if state.state_details.deferred: 526 ↛ 527line 526 didn't jump to line 527 because the condition on line 526 was never true
527 task_policy = BackgroundTaskPolicy # CoreTaskPolicy + prevent `Running` -> `Running` transition
528 elif force or task_policy is None: 528 ↛ 531line 528 didn't jump to line 531 because the condition on line 528 was always true
529 task_policy = MinimalTaskPolicy
531 orchestration_rules = task_policy.compile_transition_rules(*intended_transition) # type: ignore
532 global_rules = GlobalTaskPolicy.compile_transition_rules(*intended_transition)
534 context = TaskOrchestrationContext(
535 session=session,
536 run=run,
537 initial_state=initial_state,
538 proposed_state=state,
539 )
541 if orchestration_parameters is not None: 541 ↛ 545line 541 didn't jump to line 545 because the condition on line 541 was always true
542 context.parameters = orchestration_parameters
544 # apply orchestration rules and create the new task run state
545 async with contextlib.AsyncExitStack() as stack:
546 for rule in orchestration_rules:
547 context = await stack.enter_async_context(
548 rule(context, *intended_transition)
549 )
551 for rule in global_rules:
552 context = await stack.enter_async_context(
553 rule(context, *intended_transition)
554 )
556 await context.validate_proposed_state()
558 if context.orchestration_error is not None:
559 raise context.orchestration_error
561 result = OrchestrationResult(
562 state=context.validated_state,
563 status=context.response_status,
564 details=context.response_details,
565 )
567 return result
570async def with_system_labels_for_task_run(
571 session: AsyncSession,
572 task_run: schemas.core.TaskRun,
573) -> schemas.core.KeyValueLabels:
574 """Augment user supplied labels with system default labels for a task
575 run."""
577 client_supplied_labels = task_run.labels or {}
578 default_labels = cast(schemas.core.KeyValueLabels, {})
579 parent_labels: schemas.core.KeyValueLabels = {}
581 if task_run.flow_run_id:
582 default_labels["prefect.flow-run.id"] = str(task_run.flow_run_id)
583 flow_run = await models.flow_runs.read_flow_run(
584 session=session, flow_run_id=task_run.flow_run_id
585 )
586 parent_labels = flow_run.labels if flow_run and flow_run.labels else {}
588 return parent_labels | default_labels | client_supplied_labels