This is an automated email from the ASF dual-hosted git repository.
ferruzzi pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 86e86bd07d0 Fix missing _team_name on DagRun before some listener
calls (#70760)
86e86bd07d0 is described below
commit 86e86bd07d0c3a7a7387195b150137795bb7812e
Author: Kacper Muda <[email protected]>
AuthorDate: Thu Aug 6 18:24:13 2026 -0400
Fix missing _team_name on DagRun before some listener calls (#70760)
---
.../src/airflow/jobs/scheduler_job_runner.py | 39 ++++--
airflow-core/tests/unit/jobs/test_scheduler_job.py | 149 +++++++++++++++++++++
2 files changed, 176 insertions(+), 12 deletions(-)
diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py
b/airflow-core/src/airflow/jobs/scheduler_job_runner.py
index c6b555cf238..8b10948f7de 100644
--- a/airflow-core/src/airflow/jobs/scheduler_job_runner.py
+++ b/airflow-core/src/airflow/jobs/scheduler_job_runner.py
@@ -458,6 +458,24 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
# Ensure all requested dag_ids are in the result (with None for those
not found)
return {dag_id: self._dag_id_to_team_name.get(dag_id) for dag_id in
dag_ids}
+ def _stamp_team_names(self, dag_runs: Collection[DagRun], session:
Session) -> None:
+ """
+ Stamp ``_team_name`` on each DagRun.
+
+ Team names are resolved via ``_get_team_names_for_dag_ids``, which
caches results in
+ ``self._dag_id_to_team_name`` for the duration of the current
scheduler loop. In
+ practice this means the first call per loop issues one batched query;
subsequent calls
+ for the same dag_ids are pure dict reads with no DB round-trip.
+ """
+ if not self._multi_team:
+ return
+ if not dag_runs:
+ return
+ team_map = self._get_team_names_for_dag_ids({dr.dag_id for dr in
dag_runs}, session)
+ for dr in dag_runs:
+ if team := team_map.get(dr.dag_id):
+ dr._team_name = team
+
def _get_workload_team_name(self, workload: SchedulerWorkload, session:
Session) -> str | None:
"""
Resolve team name for a workload using the DAG > Bundle > Team
relationship chain.
@@ -1734,12 +1752,8 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
.group_by(DagRun)
)
)
- if self._multi_team and paused_runs:
- paused_dag_ids = {dr.dag_id for dr in paused_runs}
- paused_team_mapping =
self._get_team_names_for_dag_ids(paused_dag_ids, session)
- for dr in paused_runs:
- if team := paused_team_mapping.get(dr.dag_id):
- dr._team_name = team
+ # Team name should be added before listeners are called in
update_state()
+ self._stamp_team_names(paused_runs, session)
for dag_run in paused_runs:
dag = self.scheduler_dag_bag.get_dag_for_run(dag_run=dag_run,
session=session)
if dag is not None:
@@ -1996,12 +2010,8 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
)
)
- if self._multi_team and dag_runs:
- unique_dag_ids = {dr.dag_id for dr in dag_runs}
- dr_team_mapping =
self._get_team_names_for_dag_ids(unique_dag_ids, session)
- for dr in dag_runs:
- if team := dr_team_mapping.get(dr.dag_id):
- dr._team_name = team
+ # Team name should be added before listeners are called in
_schedule_all_dag_runs()
+ self._stamp_team_names(dag_runs, session)
callback_tuples = self._schedule_all_dag_runs(guard, dag_runs,
session)
@@ -2827,6 +2837,9 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
partial(self.scheduler_dag_bag.get_dag_for_run, session=session)
)
+ # Team name should be added before listeners are called in
notify_dagrun_state_changed()
+ self._stamp_team_names(dag_runs, session)
+
for dag_run in dag_runs:
dag_id = dag_run.dag_id
run_id = dag_run.run_id
@@ -2964,6 +2977,8 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin):
execute=False,
)
+ # Team name should be added before listeners are called in
notify_dagrun_state_changed()
+ self._stamp_team_names([dag_run], session)
dag_run.notify_dagrun_state_changed(msg="timed_out")
if dag_run.end_date and dag_run.start_date:
duration = dag_run.end_date - dag_run.start_date
diff --git a/airflow-core/tests/unit/jobs/test_scheduler_job.py
b/airflow-core/tests/unit/jobs/test_scheduler_job.py
index d344df97c20..440e6e2648c 100644
--- a/airflow-core/tests/unit/jobs/test_scheduler_job.py
+++ b/airflow-core/tests/unit/jobs/test_scheduler_job.py
@@ -342,6 +342,14 @@ class TestSchedulerJob:
yield
self.null_exec = None
+ @pytest.fixture
+ def team_bundle(self, testing_team, testing_dag_bundle, session):
+ team = session.merge(testing_team)
+ bundle =
session.scalar(select(DagBundleModel).where(DagBundleModel.name == "testing"))
+ bundle.teams.append(team)
+ session.flush()
+ return bundle
+
@pytest.fixture
def mock_executors(self):
mock_jwt_generator = MagicMock(spec=JWTGenerator)
@@ -9785,6 +9793,147 @@ class TestSchedulerJob:
assert call_args.kwargs["msg"] == "timed_out"
assert call_args.kwargs["dag_run"] == dag_run
+ @conf_vars({("core", "multi_team"): "true"})
+ @mock.patch("airflow.models.dagrun.get_listener_manager")
+ def test_dag_start_notifies_listener_with_team_name(
+ self, mock_get_listener_manager, dag_maker, session, team_bundle
+ ):
+ """Test that on_dag_run_running receives dag_run with _team_name
set."""
+ mock_listener_manager = MagicMock()
+ mock_get_listener_manager.return_value = mock_listener_manager
+
+ with dag_maker(dag_id="test_dag_start_team", bundle_name="testing",
session=session):
+ EmptyOperator(task_id="test_task")
+
+ dag_maker.create_dagrun(run_id="test_run", state=DagRunState.QUEUED)
+ session.commit()
+
+ mock_executor = MagicMock()
+ scheduler_job = Job()
+ self.job_runner = SchedulerJobRunner(scheduler_job,
executors=[mock_executor])
+
+ self.job_runner._start_queued_dagruns(session)
+
+ mock_listener_manager.hook.on_dag_run_running.assert_called_once()
+ call_args = mock_listener_manager.hook.on_dag_run_running.call_args
+ assert call_args.kwargs["dag_run"]._team_name == "testing"
+
+ @conf_vars({("core", "multi_team"): "true"})
+ @time_machine.travel(DEFAULT_DATE, tick=False)
+ @mock.patch("airflow.models.dagrun.get_listener_manager")
+ def test_dag_timeout_notifies_listener_with_team_name(
+ self, mock_get_listener_manager, dag_maker, session, team_bundle
+ ):
+ """Test that on_dag_run_failed receives dag_run with _team_name set
when a DAG times out."""
+ mock_listener_manager = MagicMock()
+ mock_get_listener_manager.return_value = mock_listener_manager
+
+ with dag_maker(
+ dag_id="test_dag_timeout_team",
+ bundle_name="testing",
+ session=session,
+ dagrun_timeout=timedelta(seconds=60),
+ ):
+ EmptyOperator(task_id="test_task")
+
+ dag_run = dag_maker.create_dagrun(run_id="test_run",
state=DagRunState.RUNNING)
+ # We set it to double dagrun timeout so the timeout path is taken.
+ dag_run.start_date = DEFAULT_DATE - timedelta(seconds=120)
+ session.merge(dag_run)
+ session.commit()
+
+ mock_executor = MagicMock()
+ scheduler_job = Job()
+ self.job_runner = SchedulerJobRunner(scheduler_job,
executors=[mock_executor])
+
+ self.job_runner._schedule_dag_run(dag_run, session)
+
+ mock_listener_manager.hook.on_dag_run_failed.assert_called_once()
+ call_args = mock_listener_manager.hook.on_dag_run_failed.call_args
+ assert call_args.kwargs["msg"] == "timed_out"
+ assert call_args.kwargs["dag_run"]._team_name == "testing"
+
+ @conf_vars({("core", "multi_team"): "true"})
+ @mock.patch("airflow.models.dagrun.get_listener_manager")
+ def test_dag_success_notifies_listener_with_team_name(
+ self, mock_get_listener_manager, dag_maker, session, team_bundle
+ ):
+ """Test that on_dag_run_success receives dag_run with _team_name
set."""
+ mock_listener_manager = MagicMock()
+ mock_get_listener_manager.return_value = mock_listener_manager
+
+ with dag_maker(dag_id="test_dag_success_team", bundle_name="testing",
session=session):
+ EmptyOperator(task_id="test_task")
+
+ dag_run = dag_maker.create_dagrun()
+ ti = dag_run.get_task_instance("test_task")
+ ti.set_state(TaskInstanceState.SUCCESS, session=session)
+
+ scheduler_job = Job()
+ self.job_runner = SchedulerJobRunner(scheduler_job,
executors=[MockExecutor(do_update=False)])
+
+ self.job_runner._do_scheduling(session)
+
+ mock_listener_manager.hook.on_dag_run_success.assert_called_once()
+ call_args = mock_listener_manager.hook.on_dag_run_success.call_args
+ assert call_args.kwargs["dag_run"]._team_name == "testing"
+
+ @conf_vars({("core", "multi_team"): "true"})
+ @mock.patch("airflow.models.dagrun.get_listener_manager")
+ def test_dag_failure_notifies_listener_with_team_name(
+ self, mock_get_listener_manager, dag_maker, session, team_bundle
+ ):
+ """Test that on_dag_run_failed receives dag_run with _team_name set."""
+ mock_listener_manager = MagicMock()
+ mock_get_listener_manager.return_value = mock_listener_manager
+
+ with dag_maker(dag_id="test_dag_failure_team", bundle_name="testing",
session=session):
+ EmptyOperator(task_id="test_task")
+
+ dag_run = dag_maker.create_dagrun()
+ ti = dag_run.get_task_instance("test_task")
+ ti.set_state(TaskInstanceState.FAILED, session=session)
+
+ scheduler_job = Job()
+ self.job_runner = SchedulerJobRunner(scheduler_job,
executors=[MockExecutor(do_update=False)])
+
+ self.job_runner._do_scheduling(session)
+
+ mock_listener_manager.hook.on_dag_run_failed.assert_called_once()
+ call_args = mock_listener_manager.hook.on_dag_run_failed.call_args
+ assert call_args.kwargs["dag_run"]._team_name == "testing"
+
+ @conf_vars({("core", "multi_team"): "true"})
+ @time_machine.travel(DEFAULT_DATE, tick=False)
+ @mock.patch("airflow.models.dagrun.get_listener_manager")
+ def test_dag_paused_success_notifies_listener_with_team_name(
+ self, mock_get_listener_manager, dag_maker, session, team_bundle
+ ):
+ """Test that on_dag_run_success receives dag_run with _team_name set
for paused DAGs."""
+ mock_listener_manager = MagicMock()
+ mock_get_listener_manager.return_value = mock_listener_manager
+
+ with dag_maker(dag_id="test_dag_paused_team", bundle_name="testing",
session=session) as dag:
+ EmptyOperator(task_id="test_task")
+
+ dag_run = dag_maker.create_dagrun()
+ dag_run.last_scheduling_decision = DEFAULT_DATE - timedelta(minutes=1)
+ ti = dag_run.get_task_instance("test_task")
+ ti.set_state(TaskInstanceState.SUCCESS, session=session)
+ dm = DagModel.get_dagmodel(dag.dag_id, session=session)
+ dm.is_paused = True
+ session.flush()
+
+ mock_executor = MagicMock()
+ scheduler_job = Job()
+ self.job_runner = SchedulerJobRunner(scheduler_job,
executors=[mock_executor])
+
+ self.job_runner._update_dag_run_state_for_paused_dags(session=session)
+
+ mock_listener_manager.hook.on_dag_run_success.assert_called_once()
+ call_args = mock_listener_manager.hook.on_dag_run_success.call_args
+ assert call_args.kwargs["dag_run"]._team_name == "testing"
+
@mock.patch("airflow.models.Deadline.handle_miss")
def test_process_expired_deadlines(self, mock_handle_miss, session,
dag_maker):
"""Verify all expired and unhandled deadlines (and only those) are
processed by the scheduler."""