Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/api/task_runs.py: 67%
142 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"""
2Routes for interacting with task run objects.
3"""
5import asyncio
6import datetime
7from typing import TYPE_CHECKING, Any, Dict, List, Optional
8from uuid import UUID
10from docket import Depends as DocketDepends
11from docket import Retry
12from fastapi import (
13 Body,
14 Depends,
15 HTTPException,
16 Path,
17 Response,
18 WebSocket,
19)
20from starlette.responses import JSONResponse
21from starlette.websockets import WebSocketDisconnect
23import prefect.server.api.dependencies as dependencies
24import prefect.server.models as models
25import prefect.server.schemas as schemas
26from prefect._internal.compatibility.starlette import status
27from prefect.logging import get_logger
28from prefect.server.api.run_history import run_history
29from prefect.server.database import PrefectDBInterface, provide_database_interface
30from prefect.server.orchestration import dependencies as orchestration_dependencies
31from prefect.server.orchestration.core_policy import CoreTaskPolicy
32from prefect.server.orchestration.policies import TaskRunOrchestrationPolicy
33from prefect.server.schemas.responses import (
34 OrchestrationResult,
35 TaskRunPaginationResponse,
36)
37from prefect.server.task_queue import MultiQueue, TaskQueue
38from prefect.server.utilities import subscriptions
39from prefect.server.utilities.server import PrefectRouter
40from prefect.types import DateTime
41from prefect.types._datetime import now
43if TYPE_CHECKING: 43 ↛ 44line 43 didn't jump to line 44 because the condition on line 43 was never true
44 import logging
46logger: "logging.Logger" = get_logger("server.api")
48router: PrefectRouter = PrefectRouter(prefix="/task_runs", tags=["Task Runs"])
51@router.post("/")
52async def create_task_run(
53 task_run: schemas.actions.TaskRunCreate,
54 response: Response,
55 db: PrefectDBInterface = Depends(provide_database_interface),
56 orchestration_parameters: Dict[str, Any] = Depends(
57 orchestration_dependencies.provide_task_orchestration_parameters
58 ),
59) -> schemas.core.TaskRun:
60 """
61 Create a task run. If a task run with the same flow_run_id,
62 task_key, and dynamic_key already exists, the existing task
63 run will be returned.
65 If no state is provided, the task run will be created in a PENDING state.
67 For more information, see https://docs.prefect.io/v3/concepts/tasks.
68 """
69 # hydrate the input model into a full task run / state model
70 task_run_dict = task_run.model_dump()
71 if not task_run_dict.get("id"):
72 task_run_dict.pop("id", None)
73 task_run = schemas.core.TaskRun(**task_run_dict)
75 if not task_run.state:
76 task_run.state = schemas.states.Pending()
78 right_now = now("UTC")
80 async with db.session_context(begin_transaction=True) as session:
81 model = await models.task_runs.create_task_run(
82 session=session,
83 task_run=task_run,
84 orchestration_parameters=orchestration_parameters,
85 )
87 if model.created >= right_now:
88 response.status_code = status.HTTP_201_CREATED
90 new_task_run: schemas.core.TaskRun = schemas.core.TaskRun.model_validate(model)
92 return new_task_run
95@router.patch("/{id:uuid}", status_code=status.HTTP_204_NO_CONTENT)
96async def update_task_run(
97 task_run: schemas.actions.TaskRunUpdate,
98 task_run_id: UUID = Path(..., description="The task run id", alias="id"),
99 db: PrefectDBInterface = Depends(provide_database_interface),
100) -> None:
101 """
102 Updates a task run.
103 """
104 async with db.session_context(begin_transaction=True) as session:
105 result = await models.task_runs.update_task_run(
106 session=session, task_run=task_run, task_run_id=task_run_id
107 )
108 if not result:
109 raise HTTPException(status.HTTP_404_NOT_FOUND, detail="Task run not found")
112@router.post("/count")
113async def count_task_runs(
114 db: PrefectDBInterface = Depends(provide_database_interface),
115 flows: schemas.filters.FlowFilter = None,
116 flow_runs: schemas.filters.FlowRunFilter = None,
117 task_runs: schemas.filters.TaskRunFilter = None,
118 deployments: schemas.filters.DeploymentFilter = None,
119) -> int:
120 """
121 Count task runs.
122 """
123 async with db.session_context() as session:
124 return await models.task_runs.count_task_runs(
125 session=session,
126 flow_filter=flows,
127 flow_run_filter=flow_runs,
128 task_run_filter=task_runs,
129 deployment_filter=deployments,
130 )
133@router.post("/history")
134async def task_run_history(
135 history_start: DateTime = Body(..., description="The history's start time."),
136 history_end: DateTime = Body(..., description="The history's end time."),
137 # Workaround for the fact that FastAPI does not let us configure ser_json_timedelta
138 # to represent timedeltas as floats in JSON.
139 history_interval_seconds: float = Body(
140 ...,
141 description=(
142 "The size of each history interval, in seconds. Must be at least 1 second."
143 ),
144 json_schema_extra={"format": "time-delta"},
145 ),
146 flows: schemas.filters.FlowFilter = None,
147 flow_runs: schemas.filters.FlowRunFilter = None,
148 task_runs: schemas.filters.TaskRunFilter = None,
149 deployments: schemas.filters.DeploymentFilter = None,
150 db: PrefectDBInterface = Depends(provide_database_interface),
151) -> List[schemas.responses.HistoryResponse]:
152 """
153 Query for task run history data across a given range and interval.
154 """
155 history_interval = datetime.timedelta(seconds=history_interval_seconds)
157 if history_interval < datetime.timedelta(seconds=1): 157 ↛ 163line 157 didn't jump to line 163 because the condition on line 157 was always true
158 raise HTTPException(
159 status.HTTP_422_UNPROCESSABLE_ENTITY,
160 detail="History interval must not be less than 1 second.",
161 )
163 async with db.session_context() as session:
164 return await run_history(
165 session=session,
166 run_type="task_run",
167 history_start=history_start,
168 history_end=history_end,
169 history_interval=history_interval,
170 flows=flows,
171 flow_runs=flow_runs,
172 task_runs=task_runs,
173 deployments=deployments,
174 )
177@router.get("/{id:uuid}")
178async def read_task_run(
179 task_run_id: UUID = Path(..., description="The task run id", alias="id"),
180 db: PrefectDBInterface = Depends(provide_database_interface),
181) -> schemas.core.TaskRun:
182 """
183 Get a task run by id.
184 """
185 async with db.session_context() as session:
186 task_run = await models.task_runs.read_task_run(
187 session=session, task_run_id=task_run_id
188 )
189 if not task_run:
190 raise HTTPException(status.HTTP_404_NOT_FOUND, detail="Task not found")
191 return task_run
194@router.post("/filter")
195async def read_task_runs(
196 sort: schemas.sorting.TaskRunSort = Body(schemas.sorting.TaskRunSort.ID_DESC),
197 limit: int = dependencies.LimitBody(),
198 offset: int = Body(0, ge=0),
199 flows: Optional[schemas.filters.FlowFilter] = None,
200 flow_runs: Optional[schemas.filters.FlowRunFilter] = None,
201 task_runs: Optional[schemas.filters.TaskRunFilter] = None,
202 deployments: Optional[schemas.filters.DeploymentFilter] = None,
203 db: PrefectDBInterface = Depends(provide_database_interface),
204) -> List[schemas.core.TaskRun]:
205 """
206 Query for task runs.
207 """
208 async with db.session_context() as session:
209 return await models.task_runs.read_task_runs(
210 session=session,
211 flow_filter=flows,
212 flow_run_filter=flow_runs,
213 task_run_filter=task_runs,
214 deployment_filter=deployments,
215 offset=offset,
216 limit=limit,
217 sort=sort,
218 )
221@router.post("/paginate", response_class=JSONResponse)
222async def paginate_task_runs(
223 sort: schemas.sorting.TaskRunSort = Body(schemas.sorting.TaskRunSort.ID_DESC),
224 limit: int = dependencies.LimitBody(),
225 page: int = Body(1, ge=1),
226 flows: Optional[schemas.filters.FlowFilter] = None,
227 flow_runs: Optional[schemas.filters.FlowRunFilter] = None,
228 task_runs: Optional[schemas.filters.TaskRunFilter] = None,
229 deployments: Optional[schemas.filters.DeploymentFilter] = None,
230 db: PrefectDBInterface = Depends(provide_database_interface),
231) -> TaskRunPaginationResponse:
232 """
233 Pagination query for task runs.
234 """
235 offset = (page - 1) * limit
237 async def get_runs():
238 async with db.session_context() as session:
239 return await models.task_runs.read_task_runs(
240 session=session,
241 flow_filter=flows,
242 flow_run_filter=flow_runs,
243 task_run_filter=task_runs,
244 deployment_filter=deployments,
245 offset=offset,
246 limit=limit,
247 sort=sort,
248 )
250 async def get_count():
251 async with db.session_context() as session:
252 return await models.task_runs.count_task_runs(
253 session=session,
254 flow_filter=flows,
255 flow_run_filter=flow_runs,
256 task_run_filter=task_runs,
257 deployment_filter=deployments,
258 )
260 runs, total_count = await asyncio.gather(get_runs(), get_count())
262 return TaskRunPaginationResponse.model_validate(
263 dict(
264 results=runs,
265 count=total_count,
266 limit=limit,
267 pages=(total_count + limit - 1) // limit,
268 page=page,
269 )
270 )
273@router.delete("/{id:uuid}", status_code=status.HTTP_204_NO_CONTENT)
274async def delete_task_run(
275 docket: dependencies.Docket,
276 task_run_id: UUID = Path(..., description="The task run id", alias="id"),
277 db: PrefectDBInterface = Depends(provide_database_interface),
278) -> None:
279 """
280 Delete a task run by id.
281 """
282 async with db.session_context(begin_transaction=True) as session:
283 result = await models.task_runs.delete_task_run(
284 session=session, task_run_id=task_run_id
285 )
286 if not result:
287 raise HTTPException(status.HTTP_404_NOT_FOUND, detail="Task not found")
288 await docket.add(
289 delete_task_run_logs,
290 key=f"delete_task_run_logs:{task_run_id}",
291 )(task_run_id=task_run_id)
294async def delete_task_run_logs(
295 *,
296 db: PrefectDBInterface = DocketDepends(provide_database_interface),
297 task_run_id: UUID,
298 retry: Retry = Retry(attempts=5, delay=datetime.timedelta(seconds=0.5)),
299) -> None:
300 async with db.session_context(begin_transaction=True) as session:
301 await models.logs.delete_logs(
302 session=session,
303 log_filter=schemas.filters.LogFilter(
304 task_run_id=schemas.filters.LogFilterTaskRunId(any_=[task_run_id])
305 ),
306 )
309@router.post("/{id:uuid}/set_state")
310async def set_task_run_state(
311 task_run_id: UUID = Path(..., description="The task run id", alias="id"),
312 state: schemas.actions.StateCreate = Body(..., description="The intended state."),
313 force: bool = Body(
314 False,
315 description=(
316 "If false, orchestration rules will be applied that may alter or prevent"
317 " the state transition. If True, orchestration rules are not applied."
318 ),
319 ),
320 db: PrefectDBInterface = Depends(provide_database_interface),
321 response: Response = None,
322 task_policy: TaskRunOrchestrationPolicy = Depends(
323 orchestration_dependencies.provide_task_policy
324 ),
325 orchestration_parameters: Dict[str, Any] = Depends(
326 orchestration_dependencies.provide_task_orchestration_parameters
327 ),
328) -> OrchestrationResult:
329 """Set a task run state, invoking any orchestration rules."""
331 right_now = now("UTC")
333 # create the state
334 async with db.session_context(
335 begin_transaction=True, with_for_update=True
336 ) as session:
337 orchestration_result = await models.task_runs.set_task_run_state(
338 session=session,
339 task_run_id=task_run_id,
340 state=schemas.states.State.model_validate(
341 state
342 ), # convert to a full State object
343 force=force,
344 task_policy=CoreTaskPolicy,
345 orchestration_parameters=orchestration_parameters,
346 )
348 # set the 201 if a new state was created
349 if orchestration_result.state and orchestration_result.state.timestamp >= right_now: 349 ↛ 352line 349 didn't jump to line 352 because the condition on line 349 was always true
350 response.status_code = status.HTTP_201_CREATED
351 else:
352 response.status_code = status.HTTP_200_OK
354 return orchestration_result
357@router.websocket("/subscriptions/scheduled")
358async def scheduled_task_subscription(websocket: WebSocket) -> None:
359 websocket = await subscriptions.accept_prefect_socket(websocket)
360 if not websocket:
361 return
363 try:
364 subscription = await websocket.receive_json()
365 except subscriptions.NORMAL_DISCONNECT_EXCEPTIONS:
366 return
368 if subscription.get("type") != "subscribe":
369 return await websocket.close(
370 code=4001, reason="Protocol violation: expected 'subscribe' message"
371 )
373 task_keys = subscription.get("keys", [])
374 if not task_keys:
375 return await websocket.close(
376 code=4001, reason="Protocol violation: expected 'keys' in subscribe message"
377 )
379 if not (client_id := subscription.get("client_id")):
380 return await websocket.close(
381 code=4001,
382 reason="Protocol violation: expected 'client_id' in subscribe message",
383 )
385 subscribed_queue = MultiQueue(task_keys)
387 logger.info(f"Task worker {client_id!r} subscribed to task keys {task_keys!r}")
389 while True:
390 try:
391 # observe here so that all workers with active websockets are tracked
392 await models.task_workers.observe_worker(task_keys, client_id)
393 task_run = await asyncio.wait_for(subscribed_queue.get(), timeout=1)
394 except asyncio.TimeoutError:
395 if not await subscriptions.still_connected(websocket):
396 await models.task_workers.forget_worker(client_id)
397 return
398 continue
400 try:
401 await websocket.send_json(task_run.model_dump(mode="json"))
403 acknowledgement = await websocket.receive_json()
404 ack_type = acknowledgement.get("type")
405 if ack_type != "ack":
406 if ack_type == "quit":
407 return await websocket.close()
409 raise WebSocketDisconnect(
410 code=4001, reason="Protocol violation: expected 'ack' message"
411 )
413 await models.task_workers.observe_worker([task_run.task_key], client_id)
415 except subscriptions.NORMAL_DISCONNECT_EXCEPTIONS:
416 # If sending fails or pong fails, put the task back into the retry queue
417 await asyncio.shield(TaskQueue.for_key(task_run.task_key).retry(task_run))
418 return
419 finally:
420 await models.task_workers.forget_worker(client_id)