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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
1from __future__ import annotations
3import asyncio
4from contextlib import asynccontextmanager
5from typing import TYPE_CHECKING, Any, AsyncGenerator, NoReturn, Optional
6from uuid import UUID
8import sqlalchemy as sa
9from pydantic import BaseModel
10from sqlalchemy.exc import IntegrityError
11from sqlalchemy.ext.asyncio import AsyncSession
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
37if TYPE_CHECKING: 37 ↛ 38line 37 didn't jump to line 38 because the condition on line 37 was never true
38 import logging
40 TaskRunUpsertKey = tuple[str, UUID] | tuple[str, UUID, str, str]
41 TaskRunBatchKey = tuple[frozenset[str], str]
42 TaskRunBatchByKey = tuple[TaskRunBatchKey, list[dict[str, Any]]]
44logger: "logging.Logger" = get_logger(__name__)
46DEFAULT_PERSIST_MAX_RETRIES = 5
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 )
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
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
82 now = prefect.types._datetime.now("UTC")
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 )
103 logger.debug(f"Recorded {len(task_runs)} task run state change(s)")
106def task_run_from_event(event: ReceivedEvent) -> TaskRun:
107 task_run_id = event.resource.prefect_object_id("prefect.task-run")
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")
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
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 )
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)
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 )
150 assert task_run.state is not None
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 }
159 return task_run, {
160 **task_run_attributes,
161 **denormalized_state_attributes,
162 }
165async def record_task_run_event(event: ReceivedEvent, depth: int = 0) -> None:
166 """Record a single task run event in the database.
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 )
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.
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 """
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
210async def _record_bulk_task_run_events(events: list[ReceivedEvent]) -> None:
211 if len(events) == 0:
212 return
214 now = prefect.types._datetime.now("UTC")
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 ]
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] = {}
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]
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
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)
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)
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 )
261 unique_task_runs = sorted(
262 unique_task_runs_by_group.values(),
263 key=lambda tr: _task_run_upsert_key(tr["task_run"]),
264 )
266 db = provide_database_interface()
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]))
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
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"
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]))
350 logger.debug(
351 f"Partitioned task runs into {len(batches_by_keys)} groups by update columns"
352 )
354 canonical_task_run_ids: dict[TaskRunUpsertKey, UUID] = {}
356 for (column_keys, conflict_target_name), batch in batches_by_keys:
357 update_cols = set(column_keys) - {"id", "created"}
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 ]
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)
386 logger.debug(f"Finished bulk inserting {len(batch)} task runs")
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
421 for alias_key, canonical_key in upsert_key_aliases.items():
422 canonical_task_run_ids[alias_key] = canonical_task_run_ids[canonical_key]
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
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()
438class RetryableEvent(BaseModel):
439 event: ReceivedEvent
440 persist_attempts: int = 0
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 )
453 queue: asyncio.Queue[RetryableEvent] = asyncio.Queue()
455 async def flush() -> None:
456 logger.debug(f"Persisting {queue.qsize()} events...")
458 batch: list[RetryableEvent] = []
460 while queue.qsize() > 0 and len(batch) < write_batch_size:
461 batch.append(await queue.get())
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 )
483 if dropped > 0:
484 raise
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")
500 async def message_handler(message: Message):
501 event: ReceivedEvent = ReceivedEvent.model_validate_json(message.data)
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
506 if not event.resource.get("prefect.orchestration") == "client":
507 return
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 )
516 await queue.put(RetryableEvent(event=event))
518 if queue.qsize() >= write_batch_size:
519 await flush()
521 periodic_flush = asyncio.create_task(flush_periodically())
523 try:
524 yield message_handler
525 finally:
526 periodic_flush.cancel()
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()
532class TaskRunRecorder(RunInEphemeralServers, Service):
533 """Constructs task runs and states from client-emitted events"""
535 consumer_task: asyncio.Task[None] | None = None
536 metrics_task: asyncio.Task[None] | None = None
538 @classmethod
539 def service_settings(cls) -> ServicesBaseSetting:
540 return get_current_settings().server.services.task_run_recorder
542 def __init__(self):
543 super().__init__()
544 self._started_event: Optional[asyncio.Event] = None
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
552 @started_event.setter
553 def started_event(self, value: asyncio.Event) -> None:
554 self._started_event = value
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 )
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())
575 logger.debug("TaskRunRecorder started")
576 self.started_event.set()
578 try:
579 await self.consumer_task
580 except asyncio.CancelledError:
581 pass
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")