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
« 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
20from collections.abc import Generator, Iterable
21from typing import TYPE_CHECKING, Annotated, Any
22from uuid import UUID
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
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
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
79log = structlog.get_logger(logger_name=__name__)
80grid_router = AirflowRouter(prefix="/grid", tags=["Grid"])
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
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
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
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
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 )
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))
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]
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
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()
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)
247 return [GridNodeResponse(**n) for n in merged_nodes]
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 )
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
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
401 serdag = _get_serdag(dag_bag, dag_id, dag_version_id, session)
402 if TYPE_CHECKING:
403 assert serdag
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 }
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}
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).
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.
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 """
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.
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 )
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"
527 return StreamingResponse(content=_generate(), media_type="application/x-ndjson")