Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/core_api/services/public/dag_run.py: 68%
204 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
1# Licensed to the Apache Software Foundation (ASF) under one
2# or more contributor license agreements. See the NOTICE file
3# distributed with this work for additional information
4# regarding copyright ownership. The ASF licenses this file
5# to you under the Apache License, Version 2.0 (the
6# "License"); you may not use this file except in compliance
7# with the License. You may obtain a copy of the License at
8#
9# http://www.apache.org/licenses/LICENSE-2.0
10#
11# Unless required by applicable law or agreed to in writing,
12# software distributed under the License is distributed on an
13# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14# KIND, either express or implied. See the License for the
15# specific language governing permissions and limitations
16# under the License.
18from __future__ import annotations
20import asyncio
21import itertools
22import json
23import operator
24from typing import TYPE_CHECKING, Any
26import attrs
27import structlog
28from fastapi import HTTPException, status
29from sqlalchemy import select, tuple_
30from sqlalchemy.orm import Session, joinedload
32from airflow.api.common.mark_tasks import (
33 set_dag_run_state_to_failed,
34 set_dag_run_state_to_queued,
35 set_dag_run_state_to_success,
36)
37from airflow.api_fastapi.auth.managers.models.base_user import BaseUser
38from airflow.api_fastapi.common.dagbag import (
39 DagBagDep,
40 get_dag_for_run,
41 get_latest_version_of_dag,
42 resolve_run_on_latest_version,
43)
44from airflow.api_fastapi.common.db.task_instances import eager_load_TI_and_TIH_for_validation
45from airflow.api_fastapi.core_api.datamodels.common import (
46 BulkActionNotOnExistence,
47 BulkActionResponse,
48 BulkBody,
49 BulkCreateAction,
50 BulkDeleteAction,
51 BulkUpdateAction,
52)
53from airflow.api_fastapi.core_api.datamodels.dag_run import (
54 BulkDAGRunBody,
55 ClearPartitionsBody,
56 DagRunMutableStates,
57)
58from airflow.api_fastapi.core_api.datamodels.task_instances import NewTaskResponse
59from airflow.api_fastapi.core_api.services.public.common import BulkService
60from airflow.api_fastapi.core_api.services.public.task_instances import _emit_state_listener_hooks
61from airflow.listeners.listener import get_listener_manager
62from airflow.models.dagrun import DagRun, clear_partition_runs
63from airflow.models.taskinstance import TaskInstance
64from airflow.models.xcom import XCOM_RETURN_KEY, XComModel
65from airflow.utils.session import create_session_async
66from airflow.utils.state import State, TaskInstanceState
68if TYPE_CHECKING: 68 ↛ 69line 68 didn't jump to line 69 because the condition on line 68 was never true
69 from collections.abc import AsyncGenerator, Iterator
71 from airflow.serialization.definitions.dag import SerializedDAG
73log = structlog.get_logger(__name__)
76def get_dag_run_and_dag_for_clear(
77 *,
78 session: Session,
79 dag_bag: DagBagDep,
80 dag_id: str,
81 dag_run_id: str,
82) -> tuple[DagRun, SerializedDAG]:
83 dag_run = session.scalar(
84 select(DagRun).filter_by(dag_id=dag_id, run_id=dag_run_id).options(joinedload(DagRun.dag_model))
85 )
86 if dag_run is None:
87 raise HTTPException(
88 status.HTTP_404_NOT_FOUND,
89 f"The DagRun with dag_id: `{dag_id}` and run_id: `{dag_run_id}` was not found",
90 )
91 dag = dag_bag.get_dag_for_run(dag_run, session=session)
92 if not dag: 92 ↛ 93line 92 didn't jump to line 93 because the condition on line 92 was never true
93 raise HTTPException(status.HTTP_404_NOT_FOUND, f"Dag with id {dag_id} was not found")
94 return dag_run, dag
97def dry_run_clear_dag_run(
98 *,
99 session: Session,
100 dag_bag: DagBagDep,
101 dag_id: str,
102 dag_run_id: str,
103 only_failed: bool,
104 only_new: bool,
105) -> list[Any]:
106 if only_new:
107 # ``dag.clear(only_new=True, dry_run=True)`` returns nothing when
108 # ``created_dag_version_id`` is None (e.g. LocalDagBundle), so derive new
109 # tasks from TI existence instead.
110 latest_dag = get_latest_version_of_dag(dag_bag, dag_id, session)
111 existing_task_ids = set(
112 session.scalars(
113 select(TaskInstance.task_id).where(
114 TaskInstance.dag_id == dag_id,
115 TaskInstance.run_id == dag_run_id,
116 )
117 ).all()
118 )
119 new_task_ids = sorted(set(latest_dag.task_ids) - existing_task_ids)
120 return [NewTaskResponse(task_id=task_id, task_display_name=task_id) for task_id in new_task_ids]
122 ti_query = eager_load_TI_and_TIH_for_validation(select(TaskInstance))
123 ti_query = ti_query.where(
124 TaskInstance.dag_id == dag_id,
125 TaskInstance.run_id == dag_run_id,
126 )
127 if only_failed:
128 ti_query = ti_query.where(
129 TaskInstance.state.in_([TaskInstanceState.FAILED, TaskInstanceState.UPSTREAM_FAILED])
130 )
131 return list(session.scalars(ti_query))
134def perform_clear_dag_run(
135 *,
136 session: Session,
137 dag: SerializedDAG,
138 dag_run: DagRun,
139 dag_id: str,
140 only_failed: bool,
141 only_new: bool,
142 run_on_latest_version: bool | None,
143 note: str | None,
144 user: BaseUser,
145) -> DagRun:
146 resolved_run_on_latest = resolve_run_on_latest_version(run_on_latest_version, dag_id, session)
147 dag.clear(
148 run_id=dag_run.run_id,
149 task_ids=None,
150 only_new=only_new,
151 only_failed=only_failed,
152 run_on_latest_version=resolved_run_on_latest,
153 session=session,
154 )
155 dag_run_cleared = session.scalar(select(DagRun).where(DagRun.id == dag_run.id))
156 if not dag_run_cleared: 156 ↛ 157line 156 didn't jump to line 157 because the condition on line 156 was never true
157 raise HTTPException(status.HTTP_404_NOT_FOUND, "Dag run not found after clearing")
158 if note is not None:
159 patch_dag_run_note(dag_run=dag_run_cleared, note=note, user=user)
160 return dag_run_cleared
163def clear_partition_fields(
164 *,
165 dag: SerializedDAG,
166 body: ClearPartitionsBody,
167 dag_id: str,
168 session: Session,
169) -> tuple[int, int]:
170 """
171 Reset partition_key and partition_date to None on matching runs.
173 Returns (dag_runs_cleared, task_instances_cleared).
174 """
175 return clear_partition_runs(
176 dag=dag,
177 dag_id=dag_id,
178 run_id=body.run_id,
179 partition_key=body.partition_key,
180 partition_date_start=body.partition_date_start,
181 partition_date_end=body.partition_date_end,
182 clear_tis=body.clear_task_instances,
183 dry_run=body.dry_run,
184 session=session,
185 )
188def patch_dag_run_state(
189 *,
190 dag: SerializedDAG,
191 dag_run: DagRun,
192 state: DagRunMutableStates,
193 session: Session,
194) -> None:
195 """Set a Dag Run's state (success/queued/failed), firing the matching listener hooks."""
196 if state == DagRunMutableStates.SUCCESS:
197 _, killed_tis = set_dag_run_state_to_success(
198 dag=dag, run_id=dag_run.run_id, commit=True, session=session
199 )
200 _emit_state_listener_hooks(killed_tis, TaskInstanceState.SUCCESS)
201 try:
202 if dag_run.dag is None: 202 ↛ 204line 202 didn't jump to line 204 because the condition on line 202 was always true
203 dag_run.dag = dag
204 get_listener_manager().hook.on_dag_run_success(
205 dag_run=dag_run,
206 msg=f"Dag Run's state was manually set to `{DagRunMutableStates.SUCCESS.value}`.",
207 )
208 except Exception:
209 log.exception("error calling listener")
210 elif state == DagRunMutableStates.QUEUED:
211 # TODO AIP-103: https://github.com/apache/airflow/issues/66755
212 # Handle clearing states for all task instances in a dagrun when cleared.
213 # Not notifying on queued - only notifying on RUNNING, which happens in the scheduler.
214 set_dag_run_state_to_queued(dag=dag, run_id=dag_run.run_id, commit=True, session=session)
215 elif state == DagRunMutableStates.FAILED: 215 ↛ exitline 215 didn't return from function 'patch_dag_run_state' because the condition on line 215 was always true
216 _, killed_tis = set_dag_run_state_to_failed(
217 dag=dag, run_id=dag_run.run_id, commit=True, session=session
218 )
219 _emit_state_listener_hooks(killed_tis, TaskInstanceState.FAILED)
220 try:
221 if dag_run.dag is None: 221 ↛ 223line 221 didn't jump to line 223 because the condition on line 221 was always true
222 dag_run.dag = dag
223 get_listener_manager().hook.on_dag_run_failed(
224 dag_run=dag_run,
225 msg=f"Dag Run's state was manually set to `{DagRunMutableStates.FAILED.value}`.",
226 )
227 except Exception:
228 log.exception("error calling listener")
231def patch_dag_run_note(*, dag_run: DagRun, note: str | None, user: BaseUser) -> None:
232 """Set, update, or clear a Dag Run's note. An empty note removes it so the run is left without a note."""
233 if note == "":
234 dag_run.dag_run_note = None
235 elif dag_run.dag_run_note is None:
236 dag_run.note = (note, user.get_id())
237 else:
238 dag_run.dag_run_note.content = note
239 dag_run.dag_run_note.user_id = user.get_id()
242@attrs.define
243class DagRunWaiter:
244 """Wait for the specified dag run to finish, and collect info from it."""
246 dag_id: str
247 run_id: str
248 interval: float
249 result_task_ids: list[str] | None
251 async def _get_dag_run(self) -> DagRun:
252 async with create_session_async() as session:
253 return await session.scalar(select(DagRun).filter_by(dag_id=self.dag_id, run_id=self.run_id))
255 async def _serialize_xcoms(self) -> dict[str, Any]:
256 if self.result_task_ids is None: # Return dag-author-specified results.
257 xcom_query = XComModel.get_many(
258 run_id=self.run_id,
259 key=XCOM_RETURN_KEY,
260 dag_ids=self.dag_id,
261 )
262 xcom_query = xcom_query.where(XComModel.dag_result.is_(True))
263 else: # Explicitly API user-specified results.
264 xcom_query = XComModel.get_many(
265 run_id=self.run_id,
266 key=XCOM_RETURN_KEY,
267 task_ids=self.result_task_ids,
268 dag_ids=self.dag_id,
269 )
270 # XComModel.get_many() orders XCom by timestamp. Reset this to make
271 # mapped task results stable since execution order is not guaranteed.
272 xcom_query = xcom_query.order_by(None).order_by(XComModel.task_id, XComModel.map_index)
273 async with create_session_async() as session:
274 xcom_results = (await session.scalars(xcom_query)).all()
276 def _group_xcoms(g: Iterator[XComModel | tuple[XComModel]]) -> Any:
277 entries = [row[0] if isinstance(row, tuple) else row for row in g]
278 if len(entries) == 1 and entries[0].map_index < 0: # Unpack non-mapped task xcom.
279 return entries[0].value
280 return [entry.value for entry in entries] # Task is mapped; return all xcoms in a list.
282 return {
283 task_id: _group_xcoms(g)
284 for task_id, g in itertools.groupby(xcom_results, key=operator.attrgetter("task_id"))
285 }
287 async def _serialize_response(self, dag_run: DagRun) -> str:
288 resp = {"state": dag_run.state}
289 if dag_run.state not in State.finished_dr_states:
290 return json.dumps(resp)
291 if self.result_task_ids is None or self.result_task_ids:
292 if result_xcoms := await self._serialize_xcoms():
293 resp["results"] = result_xcoms
294 return json.dumps(resp)
296 async def wait(self) -> AsyncGenerator[str, None]:
297 yield await self._serialize_response(dag_run := await self._get_dag_run())
298 yield "\n"
299 while dag_run.state not in State.finished_dr_states:
300 await asyncio.sleep(self.interval)
301 yield await self._serialize_response(dag_run := await self._get_dag_run())
302 yield "\n"
305class BulkDagRunService(BulkService[BulkDAGRunBody]):
306 """Service for handling bulk operations on Dag Runs."""
308 def __init__(
309 self,
310 session: Session,
311 request: BulkBody[BulkDAGRunBody],
312 dag_id: str,
313 dag_bag: DagBagDep,
314 user: BaseUser,
315 ):
316 super().__init__(session, request)
317 self.dag_id = dag_id
318 self.dag_bag = dag_bag
319 self.user = user
321 def handle_bulk_create(
322 self, action: BulkCreateAction[BulkDAGRunBody], results: BulkActionResponse
323 ) -> None:
324 results.errors.append(
325 {
326 "error": "Dag Runs bulk create is not supported. Use the trigger Dag Run endpoint instead.",
327 "status_code": status.HTTP_405_METHOD_NOT_ALLOWED,
328 }
329 )
331 def _resolve_entity_key(
332 self, entity: str | BulkDAGRunBody, results: BulkActionResponse
333 ) -> tuple[str, str] | None:
334 """
335 Resolve the ``(dag_id, dag_run_id)`` for an entity.
337 Records a 400 error and returns ``None`` when a wildcard ``~`` leaves the
338 dag_id unresolved. Shared by the bulk update and delete handlers.
339 """
340 if isinstance(entity, str):
341 dag_id, dag_run_id = self.dag_id, entity
342 else:
343 dag_id = entity.dag_id or self.dag_id
344 dag_run_id = entity.dag_run_id
346 if dag_id == "~" or dag_run_id == "~": 346 ↛ 347line 346 didn't jump to line 347 because the condition on line 346 was never true
347 if isinstance(entity, str):
348 error_msg = (
349 "When using wildcard in path, dag_id must be specified in BulkDAGRunBody"
350 f" object, not as string for dag_run_id: {entity}"
351 )
352 else:
353 error_msg = (
354 "When using wildcard in path, dag_id must be specified in request body for"
355 f" dag_run_id: {entity.dag_run_id}"
356 )
357 results.errors.append({"error": error_msg, "status_code": status.HTTP_400_BAD_REQUEST})
358 return None
360 return (dag_id, dag_run_id)
362 def _categorize_dag_runs(
363 self, keys: set[tuple[str, str]]
364 ) -> tuple[dict[tuple[str, str], DagRun], set[tuple[str, str]], set[tuple[str, str]]]:
365 """
366 Split the requested ``(dag_id, dag_run_id)`` keys into existing and missing ones.
368 :return: tuple of (dag_run_map, matched_keys, not_found_keys). Shared by the bulk
369 update and delete handlers.
370 """
371 dag_run_map = {
372 (dr.dag_id, dr.run_id): dr
373 for dr in self.session.scalars(
374 select(DagRun).where(tuple_(DagRun.dag_id, DagRun.run_id).in_(list(keys)))
375 )
376 }
377 matched_keys = set(dag_run_map.keys())
378 not_found_keys = keys - matched_keys
379 return dag_run_map, matched_keys, not_found_keys
381 def handle_bulk_update(
382 self, action: BulkUpdateAction[BulkDAGRunBody], results: BulkActionResponse
383 ) -> None:
384 """Bulk update Dag Runs (mark as success/failed/queued and/or set a note)."""
385 entities_by_key: dict[tuple[str, str], BulkDAGRunBody] = {}
386 for entity in action.entities:
387 if isinstance(entity, str): 387 ↛ 388line 387 didn't jump to line 388 because the condition on line 387 was never true
388 results.errors.append(
389 {
390 "error": (
391 "Bulk update requires a BulkDAGRunBody object,"
392 f" not a string for dag_run_id: {entity}"
393 ),
394 "status_code": status.HTTP_400_BAD_REQUEST,
395 }
396 )
397 continue
398 key = self._resolve_entity_key(entity, results)
399 if key is not None: 399 ↛ 386line 399 didn't jump to line 386 because the condition on line 399 was always true
400 entities_by_key[key] = entity
402 if not entities_by_key:
403 return
405 to_update_keys = set(entities_by_key.keys())
406 dag_run_map, matched_keys, not_found_keys = self._categorize_dag_runs(to_update_keys)
408 try:
409 if action.action_on_non_existence == BulkActionNotOnExistence.FAIL and not_found_keys:
410 raise HTTPException(
411 status.HTTP_404_NOT_FOUND,
412 f"The DagRuns with these identifiers: {sorted(not_found_keys)} were not found",
413 )
414 update_keys = (
415 matched_keys
416 if action.action_on_non_existence == BulkActionNotOnExistence.SKIP
417 else to_update_keys
418 )
420 for key, entity in entities_by_key.items():
421 if key not in update_keys: 421 ↛ 423line 421 didn't jump to line 423 because the condition on line 421 was always true
422 continue
423 dag_id, run_id = key
424 dag_run = dag_run_map[key]
425 if entity.state is not None:
426 dag = get_dag_for_run(self.dag_bag, dag_run, session=self.session)
427 patch_dag_run_state(dag=dag, dag_run=dag_run, state=entity.state, session=self.session)
428 if entity.note is not None:
429 patch_dag_run_note(dag_run=dag_run, note=entity.note, user=self.user)
430 results.success.append(f"{dag_id}.{run_id}")
431 except HTTPException as e:
432 results.errors.append({"error": f"{e.detail}", "status_code": e.status_code})
434 def handle_bulk_delete(
435 self, action: BulkDeleteAction[BulkDAGRunBody], results: BulkActionResponse
436 ) -> None:
437 """Bulk delete Dag Runs."""
438 to_delete_keys: set[tuple[str, str]] = set()
439 for entity in action.entities:
440 key = self._resolve_entity_key(entity, results)
441 if key is not None: 441 ↛ 439line 441 didn't jump to line 439 because the condition on line 441 was always true
442 to_delete_keys.add(key)
444 if not to_delete_keys:
445 return
447 dag_run_map, matched_keys, not_found_keys = self._categorize_dag_runs(to_delete_keys)
448 deletable_states = {s.value for s in DagRunMutableStates}
450 try:
451 if action.action_on_non_existence == BulkActionNotOnExistence.FAIL and not_found_keys:
452 raise HTTPException(
453 status.HTTP_404_NOT_FOUND,
454 f"The DagRuns with these identifiers: {sorted(not_found_keys)} were not found",
455 )
456 delete_keys = (
457 matched_keys
458 if action.action_on_non_existence == BulkActionNotOnExistence.SKIP
459 else to_delete_keys
460 )
462 for dag_id, run_id in sorted(delete_keys): 462 ↛ 463line 462 didn't jump to line 463 because the loop on line 462 never started
463 dag_run = dag_run_map[(dag_id, run_id)]
464 if dag_run.state not in deletable_states:
465 results.errors.append(
466 {
467 "error": (
468 f"The DagRun with dag_id: `{dag_id}` and run_id: `{run_id}` "
469 f"cannot be deleted in {dag_run.state} state"
470 ),
471 "status_code": status.HTTP_409_CONFLICT,
472 }
473 )
474 continue
475 self.session.delete(dag_run)
476 results.success.append(f"{dag_id}.{run_id}")
477 except HTTPException as e:
478 results.errors.append({"error": f"{e.detail}", "status_code": e.status_code})