Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 27 additions & 12 deletions airflow-core/src/airflow/jobs/scheduler_job_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
149 changes: 149 additions & 0 deletions airflow-core/tests/unit/jobs/test_scheduler_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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."""
Expand Down