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
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@
)
from airflow.api_fastapi.core_api.datamodels.task_instances import NewTaskResponse
from airflow.api_fastapi.core_api.services.public.common import BulkService
from airflow.listeners.listener import get_listener_manager
from airflow.listeners.listener import get_listener_manager_for_dag
from airflow.models.dagrun import DagRun, clear_partition_runs
from airflow.models.taskinstance import TaskInstance
from airflow.models.xcom import XCOM_RETURN_KEY, XComModel
Expand Down Expand Up @@ -195,7 +195,9 @@ def patch_dag_run_state(
if state == DagRunMutableStates.SUCCESS:
set_dag_run_state_to_success(dag=dag, run_id=dag_run.run_id, commit=True, session=session)
try:
get_listener_manager().hook.on_dag_run_success(dag_run=dag_run, msg="")
get_listener_manager_for_dag(dag_run.dag_id, session=session).hook.on_dag_run_success(
dag_run=dag_run, msg=""
)
except Exception:
log.exception("error calling listener")
elif state == DagRunMutableStates.QUEUED:
Expand All @@ -206,7 +208,9 @@ def patch_dag_run_state(
elif state == DagRunMutableStates.FAILED:
set_dag_run_state_to_failed(dag=dag, run_id=dag_run.run_id, commit=True, session=session)
try:
get_listener_manager().hook.on_dag_run_failed(dag_run=dag_run, msg="")
get_listener_manager_for_dag(dag_run.dag_id, session=session).hook.on_dag_run_failed(
dag_run=dag_run, msg=""
)
except Exception:
log.exception("error calling listener")

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@
from airflow.api_fastapi.core_api.security import GetUserDep
from airflow.api_fastapi.core_api.services.public.common import BulkService
from airflow.configuration import conf
from airflow.listeners.listener import get_listener_manager
from airflow.listeners.listener import get_listener_manager_for_dag
from airflow.models.dag import DagModel
from airflow.models.taskinstance import TaskInstance as TI
from airflow.serialization.definitions.dag import SerializedDAG
Expand Down Expand Up @@ -106,20 +106,23 @@ def _validate_patch_task_instance_body(
return body.model_dump(include=fields_to_update, by_alias=True)


def _emit_state_listener_hooks(updated_tis: list[TI], new_state: str | TaskInstanceState) -> None:
def _emit_state_listener_hooks(
updated_tis: list[TI], new_state: str | TaskInstanceState, session: Session
) -> None:
"""Fire listener hooks for the given TIs based on their new state. Listener errors are logged."""
for ti in updated_tis:
listener_manager = get_listener_manager_for_dag(ti.dag_id, session=session)
try:
if new_state == TaskInstanceState.SUCCESS:
get_listener_manager().hook.on_task_instance_success(previous_state=None, task_instance=ti)
listener_manager.hook.on_task_instance_success(previous_state=None, task_instance=ti)
elif new_state == TaskInstanceState.FAILED:
get_listener_manager().hook.on_task_instance_failed(
listener_manager.hook.on_task_instance_failed(
previous_state=None,
task_instance=ti,
error=f"TaskInstance's state was manually set to `{TaskInstanceState.FAILED}`.",
)
elif new_state == TaskInstanceState.SKIPPED:
get_listener_manager().hook.on_task_instance_skipped(previous_state=None, task_instance=ti)
listener_manager.hook.on_task_instance_skipped(previous_state=None, task_instance=ti)
except Exception:
log.exception("error calling listener")

Expand Down Expand Up @@ -265,7 +268,7 @@ def _patch_task_instance_state(
if data["new_state"] == TaskInstanceState.SUCCESS:
_clear_task_state_store_on_success(updated_tis, session)

_emit_state_listener_hooks(updated_tis, data["new_state"])
_emit_state_listener_hooks(updated_tis, data["new_state"], session)

return updated_tis

Expand Down Expand Up @@ -300,7 +303,7 @@ def _patch_task_group_state(
if data["new_state"] == TaskInstanceState.SUCCESS:
_clear_task_state_store_on_success(updated_tis, session)

_emit_state_listener_hooks(updated_tis, data["new_state"])
_emit_state_listener_hooks(updated_tis, data["new_state"], session)

return updated_tis

Expand Down
58 changes: 48 additions & 10 deletions airflow-core/src/airflow/listeners/listener.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,15 +18,33 @@
from __future__ import annotations

from functools import cache
from typing import TYPE_CHECKING

from airflow._shared.listeners.listener import ListenerManager
from airflow._shared.listeners.spec import lifecycle, taskinstance
from airflow.configuration import conf
from airflow.listeners.spec import asset, dagrun, importerrors
from airflow.plugins_manager import integrate_listener_plugins

if TYPE_CHECKING:
from sqlalchemy.orm import Session


@cache
def get_listener_manager() -> ListenerManager:
def _build_listener_manager(team_name: str | None) -> ListenerManager:
_listener_manager = ListenerManager()

_listener_manager.add_hookspecs(lifecycle)
_listener_manager.add_hookspecs(dagrun)
_listener_manager.add_hookspecs(taskinstance)
_listener_manager.add_hookspecs(asset)
_listener_manager.add_hookspecs(importerrors)

integrate_listener_plugins(_listener_manager, team_name=team_name)
return _listener_manager


def get_listener_manager(team_name: str | None = None) -> ListenerManager:
"""
Get a listener manager for Airflow core.

Expand All @@ -36,17 +54,37 @@ def get_listener_manager() -> ListenerManager:
- taskinstance: on_task_instance_running, on_task_instance_success, etc.
- asset: on_asset_created, on_asset_changed, etc.
- importerrors: on_new_dag_import_error, on_existing_dag_import_error

:param team_name: In multi-team mode, the team whose (and the global) plugin
listeners this manager should contain. ``None`` yields a manager with only
global-plugin listeners (used by team-agnostic events such as lifecycle,
asset, and import-error hooks). When multi-team mode is disabled this
argument is ignored and the manager holds every plugin's listeners.

Calls are cached per ``team_name``. The single positional delegation to
``_build_listener_manager`` guarantees that ``get_listener_manager()``,
``get_listener_manager(None)`` and ``get_listener_manager(team_name=None)`` all
return the identical (global) manager instance rather than distinct cache
entries.
"""
_listener_manager = ListenerManager()
return _build_listener_manager(team_name)

_listener_manager.add_hookspecs(lifecycle)
_listener_manager.add_hookspecs(dagrun)
_listener_manager.add_hookspecs(taskinstance)
_listener_manager.add_hookspecs(asset)
_listener_manager.add_hookspecs(importerrors)

integrate_listener_plugins(_listener_manager)
return _listener_manager
def get_listener_manager_for_dag(dag_id: str, session: Session | None = None) -> ListenerManager:
"""
Get the listener manager scoped to the team that owns ``dag_id``.

When multi-team mode is disabled this returns the single manager holding all
listeners. Otherwise it resolves the Dag's owning team and returns a manager
containing the global listeners plus that team's listeners.
"""
if not conf.getboolean("core", "multi_team"):
return get_listener_manager()

from airflow.models.dag import DagModel

team_name = DagModel.get_team_name(dag_id, session=session) if session else DagModel.get_team_name(dag_id)
return get_listener_manager(team_name)


__all__ = ["get_listener_manager", "ListenerManager"]
__all__ = ["get_listener_manager", "get_listener_manager_for_dag", "ListenerManager"]
20 changes: 15 additions & 5 deletions airflow-core/src/airflow/models/dagrun.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,16 @@
from sqlalchemy.ext.associationproxy import association_proxy
from sqlalchemy.ext.hybrid import hybrid_property
from sqlalchemy.ext.mutable import MutableDict
from sqlalchemy.orm import Mapped, declared_attr, joinedload, mapped_column, relationship, synonym, validates
from sqlalchemy.orm import (
Mapped,
declared_attr,
joinedload,
mapped_column,
object_session,
relationship,
synonym,
validates,
)
from sqlalchemy.orm.exc import StaleDataError
from sqlalchemy.sql.expression import false, select
from sqlalchemy.sql.functions import coalesce
Expand All @@ -74,7 +83,7 @@
from airflow.callbacks.callback_requests import DagCallbackRequest, DagRunContext
from airflow.configuration import conf as airflow_conf
from airflow.exceptions import AirflowException, NotMapped, TaskNotFound
from airflow.listeners.listener import get_listener_manager
from airflow.listeners.listener import get_listener_manager_for_dag
from airflow.models import Deadline, Log
from airflow.models.backfill import Backfill
from airflow.models.base import Base, StringID
Expand Down Expand Up @@ -1449,12 +1458,13 @@ def _filter_tis_and_exclude_removed(dag: SerializedDAG, tis: list[TI]) -> Iterab

def notify_dagrun_state_changed(self, msg: str):
try:
listener_manager = get_listener_manager_for_dag(self.dag_id, session=object_session(self))
if self.state == DagRunState.RUNNING:
get_listener_manager().hook.on_dag_run_running(dag_run=self, msg=msg)
listener_manager.hook.on_dag_run_running(dag_run=self, msg=msg)
elif self.state == DagRunState.SUCCESS:
get_listener_manager().hook.on_dag_run_success(dag_run=self, msg=msg)
listener_manager.hook.on_dag_run_success(dag_run=self, msg=msg)
elif self.state == DagRunState.FAILED:
get_listener_manager().hook.on_dag_run_failed(dag_run=self, msg=msg)
listener_manager.hook.on_dag_run_failed(dag_run=self, msg=msg)
except Exception:
self.log.exception("Error while calling listener")
# deliberately not notifying on QUEUED
Expand Down
4 changes: 2 additions & 2 deletions airflow-core/src/airflow/models/taskinstance.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,7 @@
from airflow.configuration import conf
from airflow.exceptions import RemovedInAirflow4Warning
from airflow.executors.workloads import BaseWorkload
from airflow.listeners.listener import get_listener_manager
from airflow.listeners.listener import get_listener_manager_for_dag
from airflow.models.asset import AssetModel
from airflow.models.base import Base, StringID, TaskInstanceDependencies
from airflow.models.dag_version import DagVersion
Expand Down Expand Up @@ -1899,7 +1899,7 @@ def fetch_handle_failure_context(
ti.state = State.UP_FOR_RETRY

try:
get_listener_manager().hook.on_task_instance_failed(
get_listener_manager_for_dag(ti.dag_id, session=session).hook.on_task_instance_failed(
previous_state=TaskInstanceState.RUNNING, task_instance=ti, error=error
)
except Exception:
Expand Down
13 changes: 11 additions & 2 deletions airflow-core/src/airflow/plugins_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -324,13 +324,22 @@ def integrate_macros_plugins() -> None:
)


def integrate_listener_plugins(listener_manager: ListenerManager) -> None:
"""Add listeners from plugins."""
def integrate_listener_plugins(listener_manager: ListenerManager, team_name: str | None = None) -> None:
"""
Add listeners from plugins to the given listener manager.

In multi-team mode, only listeners from global plugins (``team_name is None``)
and plugins belonging to ``team_name`` are registered, so a team-scoped manager
never receives another team's listeners. When multi-team mode is disabled every
plugin's listeners are registered (``team_name`` is ignored).
"""
from airflow._shared.plugins_manager import (
integrate_listener_plugins as _integrate_listener_plugins,
)

plugins, _ = _get_plugins()
if conf.getboolean("core", "multi_team"):
plugins = [plugin for plugin in plugins if plugin.team_name in (None, team_name)]
_integrate_listener_plugins(listener_manager, plugins=plugins)


Expand Down
4 changes: 2 additions & 2 deletions airflow-core/tests/unit/jobs/test_scheduler_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -9020,7 +9020,7 @@ def on_failure_callback(context):
last_ti=dag_run.get_task_instance(task_id="test_task"),
)

@mock.patch("airflow.models.dagrun.get_listener_manager")
@mock.patch("airflow.models.dagrun.get_listener_manager_for_dag")
def test_dag_start_notifies_with_started_msg(self, mock_get_listener_manager, dag_maker, session):
"""Test that notify_dagrun_state_changed is called with msg='started' when DAG starts."""
mock_listener_manager = MagicMock()
Expand All @@ -9045,7 +9045,7 @@ def test_dag_start_notifies_with_started_msg(self, mock_get_listener_manager, da
assert call_args.kwargs["msg"] == "started"
assert call_args.kwargs["dag_run"].dag_id == dag_run.dag_id

@mock.patch("airflow.models.dagrun.get_listener_manager")
@mock.patch("airflow.models.dagrun.get_listener_manager_for_dag")
def test_dag_timeout_notifies_with_timed_out_msg(self, mock_get_listener_manager, dag_maker, session):
"""Test that notify_dagrun_state_changed is called with msg='timed_out' when DAG times out."""
mock_listener_manager = MagicMock()
Expand Down
Loading
Loading