Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/core_api/routes/ui/grid.py: 20%

151 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 

20from collections.abc import Generator, Iterable 

21from typing import TYPE_CHECKING, Annotated, Any 

22from uuid import UUID 

23 

24import structlog 

25from fastapi import Depends, HTTPException, Query, status 

26from fastapi.responses import StreamingResponse 

27from sqlalchemy import exists, select 

28from sqlalchemy.orm import Session, joinedload, load_only 

29 

30from airflow.api_fastapi.auth.managers.models.resource_details import DagAccessEntity 

31from airflow.api_fastapi.common.dagbag import DagBagDep 

32from airflow.api_fastapi.common.db.common import SessionDep, paginated_select 

33from airflow.api_fastapi.common.db.dag_runs import attach_dag_versions_to_runs 

34from airflow.api_fastapi.common.parameters import ( 

35 QueryDagRunRunTypesFilter, 

36 QueryDagRunStateFilter, 

37 QueryDagRunTriggeringUserPrefixSearch, 

38 QueryDagRunTriggeringUserSearch, 

39 QueryIncludeDownstream, 

40 QueryIncludeUpstream, 

41 QueryLimit, 

42 QueryOffset, 

43 RangeFilter, 

44 SortParam, 

45 datetime_range_filter_factory, 

46) 

47from airflow.api_fastapi.common.router import AirflowRouter 

48from airflow.api_fastapi.core_api.datamodels.ui.common import ( 

49 GridNodeResponse, 

50 GridRunsResponse, 

51) 

52from airflow.api_fastapi.core_api.datamodels.ui.grid import ( 

53 GridTISummaries, 

54) 

55from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc 

56from airflow.api_fastapi.core_api.security import requires_access_dag 

57from airflow.api_fastapi.core_api.services.ui.grid import ( 

58 GridNodeAgg, 

59 _find_aggregates, 

60 _get_aggs_for_node, 

61 _merge_node_dicts, 

62) 

63from airflow.api_fastapi.core_api.services.ui.task_group import ( 

64 get_task_group_children_getter, 

65 task_group_to_dict_grid, 

66) 

67from airflow.models.dag_version import DagVersion 

68from airflow.models.dagrun import DagRun, DagRunNote 

69from airflow.models.deadline import Deadline 

70from airflow.models.serialized_dag import SerializedDagModel 

71from airflow.models.taskinstance import TaskInstance, TaskInstanceNote 

72from airflow.utils.helpers import chunks 

73from airflow.utils.session import create_session 

74 

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

76 from airflow.models.dagbag import DBDagBag 

77 from airflow.serialization.definitions.dag import SerializedDAG 

78 

79log = structlog.get_logger(logger_name=__name__) 

80grid_router = AirflowRouter(prefix="/grid", tags=["Grid"]) 

81 

82 

83def _get_latest_serdag(dag_id, session): 

84 serdag = session.scalar(SerializedDagModel.latest_item_select_object(dag_id)) 

85 if not serdag: 

86 raise HTTPException( 

87 status.HTTP_404_NOT_FOUND, 

88 f"Dag with id {dag_id} was not found", 

89 ) 

90 return serdag 

91 

92 

93def _get_serdag( 

94 dag_bag: DBDagBag, 

95 dag_id: str, 

96 dag_version_id: UUID | str | None, 

97 session: Session, 

98) -> SerializedDAG | None: 

99 """Resolve the serialized Dag for a grid TI summary via the shared (cached) ``DBDagBag``.""" 

100 if dag_version_id is not None: 

101 serdag = dag_bag.get_dag(dag_version_id, session=session) 

102 if serdag is None: 

103 log.error("No serialized dag found", dag_id=dag_id, version_id=dag_version_id) 

104 return serdag 

105 

106 # Fallback: pre-3.0 upgrade — pick the oldest DagVersion for this dag_id. 

107 oldest_version_id = session.scalar( 

108 select(DagVersion.id).where(DagVersion.dag_id == dag_id).order_by(DagVersion.id).limit(1) 

109 ) 

110 if oldest_version_id is None: 

111 return None 

112 serdag = dag_bag.get_dag(oldest_version_id, session=session) 

113 if serdag is None: 

114 log.error("No serialized dag found", dag_id=dag_id, version_id=oldest_version_id) 

115 return serdag 

116 

117 

118@grid_router.get( 

119 "/structure/{dag_id}", 

120 responses=create_openapi_http_exception_doc([status.HTTP_400_BAD_REQUEST, status.HTTP_404_NOT_FOUND]), 

121 dependencies=[ 

122 Depends(requires_access_dag(method="GET", access_entity=DagAccessEntity.TASK_INSTANCE)), 

123 Depends(requires_access_dag(method="GET", access_entity=DagAccessEntity.RUN)), 

124 ], 

125 response_model_exclude_none=True, 

126) 

127def get_dag_structure( 

128 dag_id: str, 

129 session: SessionDep, 

130 offset: QueryOffset, 

131 limit: QueryLimit, 

132 order_by: Annotated[ 

133 SortParam, 

134 Depends(SortParam(["run_after", "logical_date", "start_date", "end_date"], DagRun).dynamic_depends()), 

135 ], 

136 run_after: Annotated[RangeFilter, Depends(datetime_range_filter_factory("run_after", DagRun))], 

137 run_type: QueryDagRunRunTypesFilter, 

138 state: QueryDagRunStateFilter, 

139 triggering_user: QueryDagRunTriggeringUserSearch, 

140 triggering_user_prefix: QueryDagRunTriggeringUserPrefixSearch, 

141 include_upstream: QueryIncludeUpstream = False, 

142 include_downstream: QueryIncludeDownstream = False, 

143 depth: int | None = None, 

144 root: str | None = None, 

145) -> list[GridNodeResponse]: 

146 """Return dag structure for grid view.""" 

147 latest_serdag = _get_latest_serdag(dag_id, session) 

148 latest_dag = latest_serdag.dag 

149 latest_serdag_id = latest_serdag.id 

150 session.expunge(latest_serdag) # allow GC of serdag; only latest_dag is needed from here 

151 

152 # Apply filtering if root task is specified 

153 if root: 

154 latest_dag = latest_dag.partial_subset( 

155 task_ids=root, 

156 include_upstream=include_upstream, 

157 include_downstream=include_downstream, 

158 depth=depth, 

159 ) 

160 

161 # Retrieve, sort the previous Dag Runs 

162 base_query = select(DagRun.id).where(DagRun.dag_id == dag_id) 

163 # This comparison is to fall back to Dag timetable when no order_by is provided 

164 if order_by.value == [order_by.get_primary_key_string()]: 

165 ordering = list(latest_dag.timetable.run_ordering) 

166 order_by = SortParam( 

167 allowed_attrs=ordering, 

168 model=DagRun, 

169 ).set_value(ordering) 

170 dag_runs_select_filter, _ = paginated_select( 

171 statement=base_query, 

172 order_by=order_by, 

173 offset=offset, 

174 filters=[run_after, run_type, state, triggering_user, triggering_user_prefix], 

175 limit=limit, 

176 ) 

177 run_ids = list(session.scalars(dag_runs_select_filter)) 

178 

179 task_group_sort = get_task_group_children_getter() 

180 # Built once per render and passed down, intentionally not memoized/LRU-cached: it is a 

181 # derived view of a mutable task-group tree, so a cache would go stale with no invalidation 

182 # and would pin the whole map for the group's lifetime, fighting the streaming/expunge below. 

183 # It is released when the request returns (explicitly del-eted after the latest serdag is 

184 # merged on the main path). 

185 latest_group_dict = latest_dag.task_group.get_task_group_dict() 

186 if not run_ids: 

187 nodes = [ 

188 task_group_to_dict_grid(x, group_dict=latest_group_dict) 

189 for x in task_group_sort(latest_dag.task_group, latest_group_dict) 

190 ] 

191 return [GridNodeResponse(**n) for n in nodes] 

192 

193 # Process and merge the latest serdag first 

194 merged_nodes: list[dict[str, Any]] = [] 

195 nodes = [ 

196 task_group_to_dict_grid(x, group_dict=latest_group_dict) 

197 for x in task_group_sort(latest_dag.task_group, latest_group_dict) 

198 ] 

199 _merge_node_dicts(merged_nodes, nodes) 

200 del latest_dag, latest_group_dict 

201 

202 # we get the ids so that we can split serialization into batches and balance round trips and mem usage 

203 serdag_id_query = select(SerializedDagModel.id).where( 

204 # Even though dag_id is filtered in base_query, 

205 # adding this line here can improve the performance of this endpoint 

206 SerializedDagModel.dag_id == dag_id, 

207 SerializedDagModel.id != latest_serdag_id, 

208 SerializedDagModel.dag_version_id.in_( 

209 select(TaskInstance.dag_version_id) 

210 .join(TaskInstance.dag_run) 

211 .where( 

212 DagRun.id.in_(run_ids), 

213 ) 

214 .distinct() 

215 ), 

216 ) 

217 serdag_ids = list(session.scalars(serdag_id_query)) 

218 # Release the request session's transaction/connection before the batched work. 

219 session.close() 

220 

221 for serdag_id_batch in chunks(serdag_ids, 5): # balance memory usage and round trips 

222 with create_session(scoped=False) as batch_session: 

223 serdags = batch_session.scalars( 

224 select(SerializedDagModel).where(SerializedDagModel.id.in_(serdag_id_batch)) 

225 ).all() 

226 for serdag in serdags: 

227 batch_session.expunge(serdag) # detach so `.dag` deserializes without the session 

228 # Connection is released here; deserialize + merge this batch outside the transaction. 

229 for serdag in serdags: 

230 filtered_dag = serdag.dag 

231 # Apply the same filtering to historical Dag versions 

232 if root: 

233 filtered_dag = filtered_dag.partial_subset( 

234 task_ids=root, 

235 include_upstream=include_upstream, 

236 include_downstream=include_downstream, 

237 depth=depth, 

238 ) 

239 # Merge immediately instead of collecting all Dags in memory 

240 filtered_group_dict = filtered_dag.task_group.get_task_group_dict() 

241 nodes = [ 

242 task_group_to_dict_grid(x, group_dict=filtered_group_dict) 

243 for x in task_group_sort(filtered_dag.task_group, filtered_group_dict) 

244 ] 

245 _merge_node_dicts(merged_nodes, nodes) 

246 

247 return [GridNodeResponse(**n) for n in merged_nodes] 

248 

249 

250@grid_router.get( 

251 "/runs/{dag_id}", 

252 responses=create_openapi_http_exception_doc( 

253 [ 

254 status.HTTP_400_BAD_REQUEST, 

255 status.HTTP_404_NOT_FOUND, 

256 ] 

257 ), 

258 dependencies=[ 

259 Depends( 

260 requires_access_dag( 

261 method="GET", 

262 access_entity=DagAccessEntity.TASK_INSTANCE, 

263 ) 

264 ), 

265 Depends( 

266 requires_access_dag( 

267 method="GET", 

268 access_entity=DagAccessEntity.RUN, 

269 ) 

270 ), 

271 ], 

272 response_model_exclude_none=True, 

273) 

274def get_grid_runs( 

275 dag_id: str, 

276 session: SessionDep, 

277 offset: QueryOffset, 

278 limit: QueryLimit, 

279 order_by: Annotated[ 

280 SortParam, 

281 Depends( 

282 SortParam( 

283 [ 

284 "run_after", 

285 "logical_date", 

286 "start_date", 

287 "end_date", 

288 ], 

289 DagRun, 

290 ).dynamic_depends() 

291 ), 

292 ], 

293 run_after: Annotated[RangeFilter, Depends(datetime_range_filter_factory("run_after", DagRun))], 

294 run_type: QueryDagRunRunTypesFilter, 

295 state: QueryDagRunStateFilter, 

296 triggering_user: QueryDagRunTriggeringUserSearch, 

297 triggering_user_prefix: QueryDagRunTriggeringUserPrefixSearch, 

298) -> list[GridRunsResponse]: 

299 """Get info about a run for the grid.""" 

300 # Retrieve, sort the previous Dag Runs 

301 has_missed_deadline = ( 

302 exists() 

303 .where(Deadline.dagrun_id == DagRun.id, Deadline.missed.is_(True)) 

304 .correlate(DagRun) 

305 .label("has_missed_deadline") 

306 ) 

307 has_note_subq = ( 

308 exists() 

309 .where(DagRunNote.dag_run_id == DagRun.id, DagRunNote.content.isnot(None)) 

310 .correlate(DagRun) 

311 .label("has_note") 

312 ) 

313 base_query = ( 

314 select(DagRun, has_missed_deadline, has_note_subq) 

315 .where(DagRun.dag_id == dag_id) 

316 .options( 

317 load_only( 

318 DagRun.dag_id, 

319 DagRun.run_id, 

320 DagRun.queued_at, 

321 DagRun.start_date, 

322 DagRun.end_date, 

323 DagRun.run_after, 

324 DagRun.state, 

325 DagRun.run_type, 

326 DagRun.bundle_version, 

327 ), 

328 joinedload(DagRun.created_dag_version).joinedload(DagVersion.bundle), 

329 joinedload(DagRun.created_dag_version).joinedload(DagVersion.dag_model), 

330 ) 

331 ) 

332 

333 # This comparison is to fall back to Dag timetable when no order_by is provided 

334 if order_by.value == [order_by.get_primary_key_string()]: 

335 latest_serdag = _get_latest_serdag(dag_id, session) 

336 latest_dag = latest_serdag.dag 

337 ordering = list(latest_dag.timetable.run_ordering) 

338 order_by = SortParam( 

339 allowed_attrs=ordering, 

340 model=DagRun, 

341 ).set_value(ordering) 

342 dag_runs_select_filter, _ = paginated_select( 

343 statement=base_query, 

344 order_by=order_by, 

345 offset=offset, 

346 filters=[run_after, run_type, state, triggering_user, triggering_user_prefix], 

347 limit=limit, 

348 return_total_entries=False, 

349 ) 

350 results = session.execute(dag_runs_select_filter).all() 

351 dag_runs = [run for run, _, _ in results] 

352 attach_dag_versions_to_runs(dag_runs, session=session) 

353 grid_runs = [] 

354 for run, has_missed, has_note in results: 

355 grid_runs.append( 

356 GridRunsResponse.model_validate( 

357 { 

358 "dag_id": run.dag_id, 

359 "run_id": run.run_id, 

360 "queued_at": run.queued_at, 

361 "start_date": run.start_date, 

362 "end_date": run.end_date, 

363 "run_after": run.run_after, 

364 "state": run.state, 

365 "run_type": run.run_type, 

366 "dag_versions": run.dag_versions, 

367 "has_missed_deadline": has_missed, 

368 "has_note": has_note, 

369 } 

370 ) 

371 ) 

372 return grid_runs 

373 

374 

375def _build_ti_summaries( 

376 dag_id: str, 

377 run_id: str, 

378 task_instances: Iterable[Any], 

379 session: Session, 

380 *, 

381 dag_bag: DBDagBag, 

382) -> dict[str, Any] | None: 

383 ti_details: dict[str, GridNodeAgg] = {} 

384 dag_version_id = None 

385 for ti in task_instances: 

386 # this is a simplification - we account for structure based on the first task 

387 dag_version_id = dag_version_id or ti.dag_version_id 

388 summary = ti_details.get(ti.task_id) 

389 if summary is None: 

390 summary = ti_details[ti.task_id] = GridNodeAgg() 

391 summary.add_ti( 

392 state=ti.state, 

393 start_date=ti.start_date, 

394 end_date=ti.end_date, 

395 dag_version_number=getattr(ti, "version_number", None), 

396 has_note=bool(getattr(ti, "has_note", False)), 

397 ) 

398 if not ti_details: 

399 return None 

400 

401 serdag = _get_serdag(dag_bag, dag_id, dag_version_id, session) 

402 if TYPE_CHECKING: 

403 assert serdag 

404 

405 def get_node_summaries() -> Iterable[dict[str, Any]]: 

406 yielded_task_ids: set[str] = set() 

407 for node, _ in _find_aggregates( 

408 node=serdag.task_group, 

409 parent_node=None, 

410 ti_details=ti_details, 

411 ): 

412 if node["type"] in {"task", "mapped_task"}: 

413 yielded_task_ids.add(node["task_id"]) 

414 if node["type"] == "task": 

415 node["child_states"] = None 

416 yield node 

417 missing_task_ids = set(ti_details.keys()) - yielded_task_ids 

418 for task_id in sorted(missing_task_ids): 

419 detail = ti_details[task_id] 

420 agg = _get_aggs_for_node(detail) 

421 yield { 

422 "task_id": task_id, 

423 "task_display_name": task_id, 

424 "type": "task", 

425 "parent_id": None, 

426 **agg, 

427 "child_states": None, 

428 } 

429 

430 nodes = list(get_node_summaries()) 

431 # If a group id and a task id collide, prefer the group record 

432 group_ids = {n.get("task_id") for n in nodes if n.get("type") == "group"} 

433 filtered = [n for n in nodes if not (n.get("type") == "task" and n.get("task_id") in group_ids)] 

434 return {"run_id": run_id, "dag_id": dag_id, "task_instances": filtered} 

435 

436 

437@grid_router.get( 

438 "/ti_summaries/{dag_id}", 

439 response_class=StreamingResponse, 

440 response_model=GridTISummaries, 

441 responses={ 

442 **create_openapi_http_exception_doc( 

443 [ 

444 status.HTTP_400_BAD_REQUEST, 

445 status.HTTP_404_NOT_FOUND, 

446 ] 

447 ), 

448 200: { 

449 "content": {"application/x-ndjson": {"schema": {"type": "string"}}}, 

450 "description": "NDJSON stream — one ``GridTISummaries`` JSON object per line, one per Dag run", 

451 }, 

452 }, 

453 dependencies=[ 

454 Depends( 

455 requires_access_dag( 

456 method="GET", 

457 access_entity=DagAccessEntity.TASK_INSTANCE, 

458 ) 

459 ), 

460 Depends( 

461 requires_access_dag( 

462 method="GET", 

463 access_entity=DagAccessEntity.RUN, 

464 ) 

465 ), 

466 ], 

467) 

468def get_grid_ti_summaries_stream( 

469 dag_id: str, 

470 dag_bag: DagBagDep, 

471 run_ids: Annotated[list[str] | None, Query()] = None, 

472) -> StreamingResponse: 

473 """ 

474 Stream TI summaries for multiple Dag runs as NDJSON (one JSON line per run). 

475 

476 Each line is a serialized ``GridTISummaries`` object emitted as soon as that 

477 run's task instances have been processed, so the client can render columns 

478 progressively without waiting for all runs to complete. 

479 

480 The serialized Dag structure is served from the app-wide ``DBDagBag`` cache 

481 (keyed by ``dag_version_id``), which avoids repeated deserialization across 

482 runs of the same version *and* across requests. 

483 """ 

484 

485 def _generate() -> Generator[str, None, None]: 

486 # Each iteration opens and closes its own DB session so the connection is 

487 # released between yields. This prevents a slow client from holding a 

488 # database connection open for the entire stream duration. 

489 # See https://github.com/apache/airflow/issues/65010. 

490 

491 has_note_subq = ( 

492 exists() 

493 .where(TaskInstanceNote.ti_id == TaskInstance.id, TaskInstanceNote.content.isnot(None)) 

494 .correlate(TaskInstance) 

495 .label("has_note") 

496 ) 

497 

498 for run_id in run_ids or []: 

499 with create_session(scoped=False) as session: 

500 tis = session.execute( 

501 select( 

502 TaskInstance.task_id, 

503 TaskInstance.state, 

504 TaskInstance.dag_version_id, 

505 TaskInstance.start_date, 

506 TaskInstance.end_date, 

507 DagVersion.version_number, 

508 has_note_subq, 

509 ) 

510 .outerjoin(DagVersion, TaskInstance.dag_version_id == DagVersion.id) 

511 .where(TaskInstance.dag_id == dag_id) 

512 .where(TaskInstance.run_id == run_id) 

513 .order_by(TaskInstance.task_id) 

514 .execution_options(yield_per=1000) 

515 ) 

516 summary = _build_ti_summaries( 

517 dag_id, 

518 run_id, 

519 tis, 

520 session, 

521 dag_bag=dag_bag, 

522 ) 

523 if summary is None: 

524 continue 

525 yield GridTISummaries.model_validate(summary).model_dump_json() + "\n" 

526 

527 return StreamingResponse(content=_generate(), media_type="application/x-ndjson")