2525from fastapi import Depends , HTTPException , Query , status
2626from fastapi .responses import StreamingResponse
2727from sqlalchemy import exists , select
28- from sqlalchemy .orm import Session , joinedload , load_only , selectinload
28+ from sqlalchemy .orm import joinedload , load_only , selectinload
2929
3030from airflow .api_fastapi .auth .managers .models .resource_details import DagAccessEntity
31- from airflow .api_fastapi .common .db .common import SessionDep , _get_session , paginated_select
31+ from airflow .api_fastapi .common .db .common import SessionDep , paginated_select
3232from airflow .api_fastapi .common .parameters import (
3333 QueryDagRunRunTypesFilter ,
3434 QueryDagRunStateFilter ,
6767from airflow .models .serialized_dag import SerializedDagModel
6868from airflow .models .taskinstance import TaskInstance
6969from airflow .models .taskinstancehistory import TaskInstanceHistory
70+ from airflow .utils .session import create_session
7071
7172log = structlog .get_logger (logger_name = __name__ )
7273grid_router = AirflowRouter (prefix = "/grid" , tags = ["Grid" ])
@@ -426,7 +427,6 @@ def get_node_summaries():
426427)
427428def get_grid_ti_summaries_stream (
428429 dag_id : str ,
429- session : Annotated [Session , Depends (_get_session )],
430430 run_ids : Annotated [list [str ] | None , Query ()] = None ,
431431) -> StreamingResponse :
432432 """
@@ -441,28 +441,34 @@ def get_grid_ti_summaries_stream(
441441 """
442442
443443 def _generate () -> Generator [str , None , None ]:
444+
445+ # Each iteration opens and closes its own DB session so the connection is
446+ # released between yields. This prevents a slow client from holding a
447+ # database connection open for the entire stream duration.
448+ # See https://github.com/apache/airflow/issues/65010.
444449 serdag_cache : dict = {}
445450 for run_id in run_ids or []:
446- tis = session .execute (
447- select (
448- TaskInstance .task_id ,
449- TaskInstance .state ,
450- TaskInstance .dag_version_id ,
451- TaskInstance .start_date ,
452- TaskInstance .end_date ,
453- DagVersion .version_number ,
454- )
455- .outerjoin (DagVersion , TaskInstance .dag_version_id == DagVersion .id )
456- .where (TaskInstance .dag_id == dag_id )
457- .where (TaskInstance .run_id == run_id )
458- .order_by (TaskInstance .task_id )
459- ).all ()
460- if not tis :
461- continue
462- version_id = tis [0 ].dag_version_id
463- if version_id not in serdag_cache :
464- serdag_cache [version_id ] = _get_serdag (dag_id , version_id , session )
465- summary = _build_ti_summaries (dag_id , run_id , tis , session , serdag = serdag_cache [version_id ])
451+ with create_session (scoped = False ) as session :
452+ tis = session .execute (
453+ select (
454+ TaskInstance .task_id ,
455+ TaskInstance .state ,
456+ TaskInstance .dag_version_id ,
457+ TaskInstance .start_date ,
458+ TaskInstance .end_date ,
459+ DagVersion .version_number ,
460+ )
461+ .outerjoin (DagVersion , TaskInstance .dag_version_id == DagVersion .id )
462+ .where (TaskInstance .dag_id == dag_id )
463+ .where (TaskInstance .run_id == run_id )
464+ .order_by (TaskInstance .task_id )
465+ ).all ()
466+ if not tis :
467+ continue
468+ version_id = tis [0 ].dag_version_id
469+ if version_id not in serdag_cache :
470+ serdag_cache [version_id ] = _get_serdag (dag_id , version_id , session )
471+ summary = _build_ti_summaries (dag_id , run_id , tis , session , serdag = serdag_cache [version_id ])
466472 yield GridTISummaries .model_validate (summary ).model_dump_json () + "\n "
467473
468474 return StreamingResponse (content = _generate (), media_type = "application/x-ndjson" )
0 commit comments