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

1""" 

2Routes for interacting with task run objects. 

3""" 

4 

5import asyncio 

6import datetime 

7from typing import TYPE_CHECKING, Any, Dict, List, Optional 

8from uuid import UUID 

9 

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 

22 

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 

42 

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

44 import logging 

45 

46logger: "logging.Logger" = get_logger("server.api") 

47 

48router: PrefectRouter = PrefectRouter(prefix="/task_runs", tags=["Task Runs"]) 

49 

50 

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. 

64 

65 If no state is provided, the task run will be created in a PENDING state. 

66 

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) 

74 

75 if not task_run.state: 

76 task_run.state = schemas.states.Pending() 

77 

78 right_now = now("UTC") 

79 

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 ) 

86 

87 if model.created >= right_now: 

88 response.status_code = status.HTTP_201_CREATED 

89 

90 new_task_run: schemas.core.TaskRun = schemas.core.TaskRun.model_validate(model) 

91 

92 return new_task_run 

93 

94 

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") 

110 

111 

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 ) 

131 

132 

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) 

156 

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 ) 

162 

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 ) 

175 

176 

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 

192 

193 

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 ) 

219 

220 

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 

236 

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 ) 

249 

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 ) 

259 

260 runs, total_count = await asyncio.gather(get_runs(), get_count()) 

261 

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 ) 

271 

272 

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) 

292 

293 

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 ) 

307 

308 

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.""" 

330 

331 right_now = now("UTC") 

332 

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 ) 

347 

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 

353 

354 return orchestration_result 

355 

356 

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 

362 

363 try: 

364 subscription = await websocket.receive_json() 

365 except subscriptions.NORMAL_DISCONNECT_EXCEPTIONS: 

366 return 

367 

368 if subscription.get("type") != "subscribe": 

369 return await websocket.close( 

370 code=4001, reason="Protocol violation: expected 'subscribe' message" 

371 ) 

372 

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 ) 

378 

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 ) 

384 

385 subscribed_queue = MultiQueue(task_keys) 

386 

387 logger.info(f"Task worker {client_id!r} subscribed to task keys {task_keys!r}") 

388 

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 

399 

400 try: 

401 await websocket.send_json(task_run.model_dump(mode="json")) 

402 

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() 

408 

409 raise WebSocketDisconnect( 

410 code=4001, reason="Protocol violation: expected 'ack' message" 

411 ) 

412 

413 await models.task_workers.observe_worker([task_run.task_key], client_id) 

414 

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)