Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/services/task_run_recorder.py: 27%

281 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 02:04 +0000

1from __future__ import annotations 

2 

3import asyncio 

4from contextlib import asynccontextmanager 

5from typing import TYPE_CHECKING, Any, AsyncGenerator, NoReturn, Optional 

6from uuid import UUID 

7 

8import sqlalchemy as sa 

9from pydantic import BaseModel 

10from sqlalchemy.exc import IntegrityError 

11from sqlalchemy.ext.asyncio import AsyncSession 

12 

13import prefect.types._datetime 

14from prefect.logging import get_logger 

15from prefect.server.database import ( 

16 PrefectDBInterface, 

17 db_injector, 

18 provide_database_interface, 

19) 

20from prefect.server.events.schemas.events import ReceivedEvent 

21from prefect.server.schemas.core import TaskRun 

22from prefect.server.schemas.states import State 

23from prefect.server.services.base import RunInEphemeralServers, Service 

24from prefect.server.utilities.messaging import ( 

25 Consumer, 

26 Message, 

27 MessageHandler, 

28 create_consumer, 

29) 

30from prefect.server.utilities.messaging._consumer_names import ( 

31 generate_unique_consumer_name, 

32) 

33from prefect.server.utilities.messaging.memory import log_metrics_periodically 

34from prefect.settings.context import get_current_settings 

35from prefect.settings.models.server.services import ServicesBaseSetting 

36 

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

38 import logging 

39 

40 TaskRunUpsertKey = tuple[str, UUID] | tuple[str, UUID, str, str] 

41 TaskRunBatchKey = tuple[frozenset[str], str] 

42 TaskRunBatchByKey = tuple[TaskRunBatchKey, list[dict[str, Any]]] 

43 

44logger: "logging.Logger" = get_logger(__name__) 

45 

46DEFAULT_PERSIST_MAX_RETRIES = 5 

47 

48 

49def _task_run_upsert_key(task_run: TaskRun) -> TaskRunUpsertKey: 

50 if task_run.flow_run_id is None: 

51 return ("id", task_run.id) 

52 return ( 

53 "natural-key", 

54 task_run.flow_run_id, 

55 task_run.task_key, 

56 task_run.dynamic_key, 

57 ) 

58 

59 

60def _task_run_conflict_keys(task_run: TaskRun) -> list[TaskRunUpsertKey]: 

61 keys: list[TaskRunUpsertKey] = [("id", task_run.id)] 

62 if task_run.flow_run_id is not None: 

63 keys.append( 

64 ( 

65 "natural-key", 

66 task_run.flow_run_id, 

67 task_run.task_key, 

68 task_run.dynamic_key, 

69 ) 

70 ) 

71 return keys 

72 

73 

74@db_injector 

75async def _insert_task_run_states( 

76 db: PrefectDBInterface, session: AsyncSession, task_runs: list[TaskRun] 

77): 

78 if TYPE_CHECKING: 

79 for task_run in task_runs: 

80 assert task_run.state is not None 

81 

82 now = prefect.types._datetime.now("UTC") 

83 

84 await session.execute( 

85 db.queries.insert(db.TaskRunState) 

86 .values( 

87 [ 

88 { 

89 "created": now, 

90 "task_run_id": task_run.id, 

91 **task_run.state.model_dump(), 

92 } 

93 for task_run in task_runs 

94 ] 

95 ) 

96 .on_conflict_do_nothing( 

97 index_elements=[ 

98 "id", 

99 ] 

100 ) 

101 ) 

102 

103 logger.debug(f"Recorded {len(task_runs)} task run state change(s)") 

104 

105 

106def task_run_from_event(event: ReceivedEvent) -> TaskRun: 

107 task_run_id = event.resource.prefect_object_id("prefect.task-run") 

108 

109 flow_run_id: Optional[UUID] = None 

110 if flow_run_resource := event.resource_in_role.get("flow-run"): 

111 flow_run_id = flow_run_resource.prefect_object_id("prefect.flow-run") 

112 

113 state: State = State.model_validate( 

114 { 

115 "id": event.id, 

116 "timestamp": event.occurred, 

117 **event.payload["validated_state"], 

118 } 

119 ) 

120 state.state_details.task_run_id = task_run_id 

121 state.state_details.flow_run_id = flow_run_id 

122 

123 return TaskRun.model_validate( 

124 { 

125 "id": task_run_id, 

126 "flow_run_id": flow_run_id, 

127 "state_id": state.id, 

128 "state": state, 

129 **event.payload["task_run"], 

130 } 

131 ) 

132 

133 

134def db_recordable_task_run_from_event( 

135 event: ReceivedEvent, 

136) -> tuple[TaskRun, dict[str, Any]]: 

137 task_run: TaskRun = task_run_from_event(event) 

138 

139 task_run_attributes = task_run.model_dump_for_orm( 

140 exclude={ 

141 "state_id", 

142 "state", 

143 "created", 

144 "estimated_run_time", 

145 "estimated_start_time_delta", 

146 }, 

147 exclude_unset=True, 

148 ) 

149 

150 assert task_run.state is not None 

151 

152 denormalized_state_attributes = { 

153 "state_id": task_run.state.id, 

154 "state_type": task_run.state.type, 

155 "state_name": task_run.state.name, 

156 "state_timestamp": task_run.state.timestamp, 

157 } 

158 

159 return task_run, { 

160 **task_run_attributes, 

161 **denormalized_state_attributes, 

162 } 

163 

164 

165async def record_task_run_event(event: ReceivedEvent, depth: int = 0) -> None: 

166 """Record a single task run event in the database. 

167 

168 Delegates to `record_bulk_task_run_events`, which already retries once on 

169 `IntegrityError` to recover from TOCTOU races against concurrent recorders. 

170 Any `IntegrityError` that survives the retry is treated as an unrecoverable 

171 duplicate and the event is discarded. 

172 """ 

173 try: 

174 await record_bulk_task_run_events([event]) 

175 except IntegrityError: 

176 logger.warning( 

177 "Duplicate task_run, discarding event %s", 

178 event.id, 

179 exc_info=True, 

180 ) 

181 

182 

183async def record_bulk_task_run_events(events: list[ReceivedEvent]) -> None: 

184 """Record multiple task run events in the database, taking advantage of bulk inserts. 

185 

186 Retries once on `IntegrityError` to handle TOCTOU races between concurrent 

187 recorder instances: when two batches reference the same `task_run.id` with 

188 different natural keys, one batch's existence-check SELECT may run before 

189 the other batch's INSERT commits. The retry re-runs the SELECT in a fresh 

190 session so the conflict target is chosen against the now-visible row. 

191 """ 

192 

193 max_attempts = 2 

194 for attempt in range(1, max_attempts + 1): 

195 try: 

196 await _record_bulk_task_run_events(events) 

197 return 

198 except IntegrityError: 

199 if attempt < max_attempts: 

200 logger.info( 

201 "Retrying bulk task_run upsert after IntegrityError" 

202 " (attempt %s/%s)", 

203 attempt, 

204 max_attempts, 

205 ) 

206 continue 

207 raise 

208 

209 

210async def _record_bulk_task_run_events(events: list[ReceivedEvent]) -> None: 

211 if len(events) == 0: 

212 return 

213 

214 now = prefect.types._datetime.now("UTC") 

215 

216 all_task_runs = [ 

217 {"task_run": task_run, "task_run_dict": task_run_dict, "event": event} 

218 for event in events 

219 for task_run, task_run_dict in [db_recordable_task_run_from_event(event)] 

220 ] 

221 

222 # Drop duplicate task run rows, keep the one with the latest state_timestamp. 

223 # A single bulk flush can contain events that collide on either id or natural 

224 # key, so coalesce connected conflicts before choosing the ON CONFLICT target. 

225 all_task_runs.sort(key=lambda tr: tr["task_run"].state.timestamp) 

226 parent: dict[TaskRunUpsertKey, TaskRunUpsertKey] = {} 

227 

228 def find(key: TaskRunUpsertKey) -> TaskRunUpsertKey: 

229 parent.setdefault(key, key) 

230 if parent[key] != key: 

231 parent[key] = find(parent[key]) 

232 return parent[key] 

233 

234 def union(left: TaskRunUpsertKey, right: TaskRunUpsertKey) -> None: 

235 left_root = find(left) 

236 right_root = find(right) 

237 if left_root != right_root: 

238 parent[right_root] = left_root 

239 

240 for tr in all_task_runs: 

241 conflict_keys = _task_run_conflict_keys(tr["task_run"]) 

242 for conflict_key in conflict_keys[1:]: 

243 union(conflict_keys[0], conflict_key) 

244 

245 unique_task_runs_by_group: dict[TaskRunUpsertKey, dict[str, Any]] = {} 

246 conflict_keys_by_group: dict[TaskRunUpsertKey, set[TaskRunUpsertKey]] = {} 

247 upsert_key_aliases: dict[TaskRunUpsertKey, TaskRunUpsertKey] = {} 

248 for tr in all_task_runs: 

249 conflict_keys = _task_run_conflict_keys(tr["task_run"]) 

250 conflict_group = find(conflict_keys[0]) 

251 tr["conflict_group"] = conflict_group 

252 unique_task_runs_by_group[conflict_group] = tr 

253 conflict_keys_by_group.setdefault(conflict_group, set()).update(conflict_keys) 

254 

255 for tr in all_task_runs: 

256 conflict_group = tr["conflict_group"] 

257 upsert_key_aliases[_task_run_upsert_key(tr["task_run"])] = _task_run_upsert_key( 

258 unique_task_runs_by_group[conflict_group]["task_run"] 

259 ) 

260 

261 unique_task_runs = sorted( 

262 unique_task_runs_by_group.values(), 

263 key=lambda tr: _task_run_upsert_key(tr["task_run"]), 

264 ) 

265 

266 db = provide_database_interface() 

267 

268 task_run_ids: list[UUID] = [] 

269 natural_keys: list[tuple[UUID, str, str]] = [] 

270 for conflict_keys in conflict_keys_by_group.values(): 

271 for conflict_key in conflict_keys: 

272 if len(conflict_key) == 2: 

273 task_run_ids.append(conflict_key[1]) 

274 else: 

275 natural_keys.append((conflict_key[1], conflict_key[2], conflict_key[3])) 

276 

277 async with db.session_context() as session: 

278 existing_task_run_ids: set[UUID] = set() 

279 existing_natural_keys: set[tuple[UUID, str, str]] = set() 

280 existing_task_run_ids_by_key: dict[TaskRunUpsertKey, UUID] = {} 

281 if task_run_ids or natural_keys: 

282 conditions = [] 

283 if task_run_ids: 

284 conditions.append(db.TaskRun.id.in_(task_run_ids)) 

285 if natural_keys: 

286 conditions.append( 

287 sa.tuple_( 

288 db.TaskRun.flow_run_id, 

289 db.TaskRun.task_key, 

290 db.TaskRun.dynamic_key, 

291 ).in_(natural_keys) 

292 ) 

293 result = await session.execute( 

294 sa.select( 

295 db.TaskRun.id, 

296 db.TaskRun.flow_run_id, 

297 db.TaskRun.task_key, 

298 db.TaskRun.dynamic_key, 

299 ).where(sa.or_(*conditions)) 

300 ) 

301 for task_run_id, flow_run_id, task_key, dynamic_key in result.all(): 

302 existing_task_run_ids.add(task_run_id) 

303 existing_task_run_ids_by_key[("id", task_run_id)] = task_run_id 

304 if flow_run_id is not None: 

305 existing_natural_keys.add((flow_run_id, task_key, dynamic_key)) 

306 existing_task_run_ids_by_key[ 

307 ("natural-key", flow_run_id, task_key, dynamic_key) 

308 ] = task_run_id 

309 

310 def conflict_target(tr: dict[str, Any]) -> str: 

311 task_run = tr["task_run"] 

312 natural_key = ( 

313 task_run.flow_run_id, 

314 task_run.task_key, 

315 task_run.dynamic_key, 

316 ) 

317 if ( 

318 task_run.flow_run_id is not None 

319 and natural_key in existing_natural_keys 

320 ): 

321 return "natural-key" 

322 if task_run.id in existing_task_run_ids: 

323 return "id" 

324 for conflict_key in sorted( 

325 conflict_keys_by_group[tr["conflict_group"]], key=str 

326 ): 

327 if conflict_key in existing_task_run_ids_by_key: 

328 canonical_task_run_id = existing_task_run_ids_by_key[conflict_key] 

329 task_run.id = canonical_task_run_id 

330 tr["task_run_dict"]["id"] = canonical_task_run_id 

331 return "id" 

332 if task_run.flow_run_id is not None: 

333 return "natural-key" 

334 return "id" 

335 

336 # Batch contiguous rows by keys to avoid column mismatches during bulk insert. 

337 # Keeping only contiguous runs preserves the global conflict-key sort order 

338 # established above for deterministic lock acquisition. 

339 batches_by_keys: list[TaskRunBatchByKey] = [] 

340 for tr in unique_task_runs: 

341 key_signature = ( 

342 frozenset(tr["task_run_dict"].keys()), 

343 conflict_target(tr), 

344 ) 

345 if batches_by_keys and batches_by_keys[-1][0] == key_signature: 

346 batches_by_keys[-1][1].append(tr) 

347 else: 

348 batches_by_keys.append((key_signature, [tr])) 

349 

350 logger.debug( 

351 f"Partitioned task runs into {len(batches_by_keys)} groups by update columns" 

352 ) 

353 

354 canonical_task_run_ids: dict[TaskRunUpsertKey, UUID] = {} 

355 

356 for (column_keys, conflict_target_name), batch in batches_by_keys: 

357 update_cols = set(column_keys) - {"id", "created"} 

358 

359 logger.debug(f"Preparing to bulk insert {len(batch)} task runs") 

360 to_insert = [ 

361 tr["task_run_dict"] | {"created": now, "updated": now} for tr in batch 

362 ] 

363 

364 insert_statement = db.queries.insert(db.TaskRun).values(to_insert) 

365 index_elements = ( 

366 db.orm.task_run_unique_upsert_columns 

367 if conflict_target_name == "natural-key" 

368 else ["id"] 

369 ) 

370 upsert_statement = insert_statement.on_conflict_do_update( 

371 index_elements=index_elements, 

372 set_={ 

373 # See https://www.postgresql.org/docs/current/sql-insert.html for details on excluded. 

374 # Idea is excluded.x references the proposed insertion value for column x. 

375 **{ 

376 col.name: getattr(insert_statement.excluded, col.name) 

377 for col in insert_statement.excluded 

378 if col.name in (update_cols | {"updated"}) - {"id"} 

379 }, 

380 }, 

381 where=db.TaskRun.state_timestamp 

382 < insert_statement.excluded.state_timestamp, 

383 ) 

384 await session.execute(upsert_statement) 

385 

386 logger.debug(f"Finished bulk inserting {len(batch)} task runs") 

387 

388 if conflict_target_name == "natural-key": 

389 natural_keys = [ 

390 ( 

391 tr["task_run"].flow_run_id, 

392 tr["task_run"].task_key, 

393 tr["task_run"].dynamic_key, 

394 ) 

395 for tr in batch 

396 ] 

397 result = await session.execute( 

398 sa.select( 

399 db.TaskRun.flow_run_id, 

400 db.TaskRun.task_key, 

401 db.TaskRun.dynamic_key, 

402 db.TaskRun.id, 

403 ).where( 

404 sa.tuple_( 

405 db.TaskRun.flow_run_id, 

406 db.TaskRun.task_key, 

407 db.TaskRun.dynamic_key, 

408 ).in_(natural_keys) 

409 ) 

410 ) 

411 for flow_run_id, task_key, dynamic_key, task_run_id in result.all(): 

412 canonical_task_run_ids[ 

413 ("natural-key", flow_run_id, task_key, dynamic_key) 

414 ] = task_run_id 

415 else: 

416 for tr in batch: 

417 canonical_task_run_ids[_task_run_upsert_key(tr["task_run"])] = tr[ 

418 "task_run" 

419 ].id 

420 

421 for alias_key, canonical_key in upsert_key_aliases.items(): 

422 canonical_task_run_ids[alias_key] = canonical_task_run_ids[canonical_key] 

423 

424 for tr in all_task_runs: 

425 task_run = tr["task_run"] 

426 canonical_task_run_id = canonical_task_run_ids[ 

427 _task_run_upsert_key(task_run) 

428 ] 

429 task_run.id = canonical_task_run_id 

430 if task_run.state is not None: 

431 task_run.state.state_details.task_run_id = canonical_task_run_id 

432 

433 # Insert all task run states - we only coalesce task run updates, not states 

434 await _insert_task_run_states(session, [tr["task_run"] for tr in all_task_runs]) 

435 await session.commit() 

436 

437 

438class RetryableEvent(BaseModel): 

439 event: ReceivedEvent 

440 persist_attempts: int = 0 

441 

442 

443@asynccontextmanager 

444async def consumer( 

445 write_batch_size: int, 

446 flush_every: int, 

447 max_persist_retries: int = DEFAULT_PERSIST_MAX_RETRIES, 

448) -> AsyncGenerator[MessageHandler, None]: 

449 logger.info( 

450 f"Creating TaskRunRecorder consumer with batch size {write_batch_size} and flush every {flush_every} seconds" 

451 ) 

452 

453 queue: asyncio.Queue[RetryableEvent] = asyncio.Queue() 

454 

455 async def flush() -> None: 

456 logger.debug(f"Persisting {queue.qsize()} events...") 

457 

458 batch: list[RetryableEvent] = [] 

459 

460 while queue.qsize() > 0 and len(batch) < write_batch_size: 

461 batch.append(await queue.get()) 

462 

463 try: 

464 await record_bulk_task_run_events([item.event for item in batch]) 

465 except Exception: 

466 dropped = 0 

467 to_retry = 0 

468 for item in batch: 

469 item.persist_attempts += 1 

470 if item.persist_attempts <= max_persist_retries: 

471 to_retry += 1 

472 await queue.put(item) 

473 else: 

474 dropped += 1 

475 logger.error( 

476 f"Dropping event {item.event.id} after {item.persist_attempts} failed attempts" 

477 ) 

478 logger.error( 

479 f"Error flushing {len(batch)} events ({to_retry} to retry, {dropped} dropped)", 

480 exc_info=True, 

481 ) 

482 

483 if dropped > 0: 

484 raise 

485 

486 async def flush_periodically(): 

487 while True: 

488 try: 

489 await asyncio.sleep(flush_every) 

490 if queue.qsize(): 490 ↛ 491line 490 didn't jump to line 491 because the condition on line 490 was never true

491 await flush() 

492 except asyncio.CancelledError: 

493 return 

494 except Exception: 

495 # flush() re-raises when events are dropped; this task is never 

496 # awaited, so letting that propagate would kill periodic 

497 # flushing silently and strand queued events (issue #21057) 

498 logger.exception("Error during periodic flush; continuing") 

499 

500 async def message_handler(message: Message): 

501 event: ReceivedEvent = ReceivedEvent.model_validate_json(message.data) 

502 

503 if not event.event.startswith("prefect.task-run"): 503 ↛ 506line 503 didn't jump to line 506 because the condition on line 503 was always true

504 return 

505 

506 if not event.resource.get("prefect.orchestration") == "client": 

507 return 

508 

509 logger.debug( 

510 "Received event: %s with id: %s for resource: %s", 

511 event.event, 

512 event.id, 

513 event.resource.get("prefect.resource.id"), 

514 ) 

515 

516 await queue.put(RetryableEvent(event=event)) 

517 

518 if queue.qsize() >= write_batch_size: 

519 await flush() 

520 

521 periodic_flush = asyncio.create_task(flush_periodically()) 

522 

523 try: 

524 yield message_handler 

525 finally: 

526 periodic_flush.cancel() 

527 

528 if queue.qsize(): 528 ↛ 529line 528 didn't jump to line 529 because the condition on line 528 was never true

529 await flush() 

530 

531 

532class TaskRunRecorder(RunInEphemeralServers, Service): 

533 """Constructs task runs and states from client-emitted events""" 

534 

535 consumer_task: asyncio.Task[None] | None = None 

536 metrics_task: asyncio.Task[None] | None = None 

537 

538 @classmethod 

539 def service_settings(cls) -> ServicesBaseSetting: 

540 return get_current_settings().server.services.task_run_recorder 

541 

542 def __init__(self): 

543 super().__init__() 

544 self._started_event: Optional[asyncio.Event] = None 

545 

546 @property 

547 def started_event(self) -> asyncio.Event: 

548 if self._started_event is None: 548 ↛ 550line 548 didn't jump to line 550 because the condition on line 548 was always true

549 self._started_event = asyncio.Event() 

550 return self._started_event 

551 

552 @started_event.setter 

553 def started_event(self, value: asyncio.Event) -> None: 

554 self._started_event = value 

555 

556 async def start( 

557 self, max_persist_retries: int = DEFAULT_PERSIST_MAX_RETRIES 

558 ) -> NoReturn: 

559 assert self.consumer_task is None, "TaskRunRecorder already started" 

560 self.consumer: Consumer = create_consumer( 

561 "events", 

562 group="task-run-recorder", 

563 name=generate_unique_consumer_name("task-run-recorder"), 

564 read_batch_size=self.service_settings().read_batch_size, 

565 ) 

566 

567 async with consumer( 

568 write_batch_size=self.service_settings().batch_size, 

569 flush_every=int(self.service_settings().flush_interval), 

570 max_persist_retries=max_persist_retries, 

571 ) as handler: 

572 self.consumer_task = asyncio.create_task(self.consumer.run(handler)) 

573 self.metrics_task = asyncio.create_task(log_metrics_periodically()) 

574 

575 logger.debug("TaskRunRecorder started") 

576 self.started_event.set() 

577 

578 try: 

579 await self.consumer_task 

580 except asyncio.CancelledError: 

581 pass 

582 

583 async def stop(self) -> None: 

584 assert self.consumer_task is not None, "Logger not started" 

585 self.consumer_task.cancel() 

586 if self.metrics_task: 586 ↛ 588line 586 didn't jump to line 588 because the condition on line 586 was always true

587 self.metrics_task.cancel() 

588 try: 

589 await self.consumer_task 

590 if self.metrics_task: 

591 await self.metrics_task 

592 except asyncio.CancelledError: 

593 pass 

594 finally: 

595 self.consumer_task = None 

596 self.metrics_task = None 

597 logger.debug("TaskRunRecorder stopped")