diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py b/airflow-core/src/airflow/jobs/scheduler_job_runner.py index b6a29995d5937..6e6d5dd78e7d3 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 @@ def _get_team_names_for_dag_ids( # 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 @@ def _update_dag_run_state_for_paused_dags(self, *, session: Session = NEW_SESSIO .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 @@ def _do_scheduling(self, session: Session) -> int: ) ) - 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) @@ -2844,6 +2854,9 @@ def _update_state(dag: SerializedDAG, dag_run: DagRun): 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 @@ -2981,6 +2994,8 @@ def _schedule_dag_run( 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 53ea5bc933f3f..824074057a352 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 @@ def set_instance_attrs(self) -> Generator: 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) @@ -9635,6 +9643,147 @@ def test_dag_timeout_notifies_with_timed_out_msg(self, mock_get_listener_manager 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."""