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

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. 

17 

18from __future__ import annotations 

19 

20import asyncio 

21import itertools 

22import json 

23import operator 

24from typing import TYPE_CHECKING, Any 

25 

26import attrs 

27import structlog 

28from fastapi import HTTPException, status 

29from sqlalchemy import select, tuple_ 

30from sqlalchemy.orm import Session, joinedload 

31 

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 

67 

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 

70 

71 from airflow.serialization.definitions.dag import SerializedDAG 

72 

73log = structlog.get_logger(__name__) 

74 

75 

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 

95 

96 

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] 

121 

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

132 

133 

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 

161 

162 

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. 

172 

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 ) 

186 

187 

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

229 

230 

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

240 

241 

242@attrs.define 

243class DagRunWaiter: 

244 """Wait for the specified dag run to finish, and collect info from it.""" 

245 

246 dag_id: str 

247 run_id: str 

248 interval: float 

249 result_task_ids: list[str] | None 

250 

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

254 

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

275 

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. 

281 

282 return { 

283 task_id: _group_xcoms(g) 

284 for task_id, g in itertools.groupby(xcom_results, key=operator.attrgetter("task_id")) 

285 } 

286 

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) 

295 

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" 

303 

304 

305class BulkDagRunService(BulkService[BulkDAGRunBody]): 

306 """Service for handling bulk operations on Dag Runs.""" 

307 

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 

320 

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 ) 

330 

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. 

336 

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 

345 

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 

359 

360 return (dag_id, dag_run_id) 

361 

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. 

367 

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 

380 

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 

401 

402 if not entities_by_key: 

403 return 

404 

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) 

407 

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 ) 

419 

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

433 

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) 

443 

444 if not to_delete_keys: 

445 return 

446 

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} 

449 

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 ) 

461 

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