From 3f7e6c1205c83881c59c5a846cdb00d1d05e288b Mon Sep 17 00:00:00 2001 From: Sameer Mesiah Date: Sun, 28 Jun 2026 20:11:45 +0100 Subject: [PATCH 1/5] Add the metric tag to OpenLineage emit and extraction metrics when multi-team is enabled on Airflow 3.1+. The team name is propagated through the Airflow run facet and included in ol.emit.attempts, ol.emit.failed, ol.extract, and ol.event.size. Tests are updated accordingly. --- .../providers/openlineage/plugins/adapter.py | 26 +- .../providers/openlineage/plugins/listener.py | 112 +++++- .../providers/openlineage/utils/utils.py | 25 +- .../providers/openlineage/version_compat.py | 3 +- .../unit/openlineage/plugins/test_adapter.py | 240 ++++++++++++- .../unit/openlineage/plugins/test_listener.py | 319 +++++++++++++++++- .../unit/openlineage/utils/test_utils.py | 55 +++ 7 files changed, 750 insertions(+), 30 deletions(-) diff --git a/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py b/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py index ec81e232fc8cf..617395a989d06 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py +++ b/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py @@ -18,7 +18,7 @@ import os import traceback -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Any, Literal, cast import yaml from openlineage.client import OpenLineageClient, set_producer @@ -50,12 +50,14 @@ get_dag_job_dependency_facet, get_processing_engine_facet, ) +from airflow.utils.helpers import prune_dict from airflow.utils.log.logging_mixin import LoggingMixin if TYPE_CHECKING: from datetime import datetime from airflow.providers.openlineage.extractors import OperatorLineage + from airflow.providers.openlineage.plugins.facets import AirflowRunFacet from airflow.sdk.execution_time.secrets_masker import SecretsMasker, _secrets_masker from airflow.utils.state import DagRunState else: @@ -172,10 +174,24 @@ def emit(self, event: RunEvent): event_type = event.eventType.value.lower() if event.eventType else "" transport_type = f"{self._client.transport.kind}".lower() + team_name = None + + facets = event.run.facets or {} + airflow_facet = cast("AirflowRunFacet | None", facets.get("airflow")) + + if airflow_facet: + team_name = airflow_facet.dagRun.get("dag_team_name") + try: with Stats.timer( "ol.emit.attempts", - tags={"event_type": event_type, "transport_type": transport_type}, + tags=prune_dict( + { + "event_type": event_type, + "transport_type": transport_type, + "team_name": team_name, + } + ), ): self._client.emit(redacted_event) self.log.info( @@ -184,7 +200,11 @@ def emit(self, event: RunEvent): event.run.runId, ) except Exception as e: - Stats.incr("ol.emit.failed") + Stats.incr( + "ol.emit.failed", + tags=prune_dict({"team_name": team_name}), + ) + self.log.warning( "Failed to emit OpenLineage `%s` event of id `%s` with the following exception: `%s`", event_type.upper(), diff --git a/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py b/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py index 4fa8c657126b6..38f0f7ad060e2 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py +++ b/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py @@ -41,6 +41,7 @@ from airflow.providers.openlineage.utils.utils import ( AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_2_PLUS, + DagRunInfo, get_airflow_dag_run_facet, get_airflow_debug_facet, get_airflow_job_facet, @@ -57,6 +58,7 @@ print_warning, ) from airflow.settings import configure_orm +from airflow.utils.helpers import prune_dict from airflow.utils.state import TaskInstanceState if TYPE_CHECKING: @@ -268,9 +270,19 @@ def on_running(): if not doc: doc, doc_type = get_dag_documentation(dag) + team_name = None + team_name = DagRunInfo.team_name(dagrun) + if controls.extract_operator_metadata: with Stats.timer( - "ol.extract", tags={"event_type": event_type, "operator_name": operator_name} + "ol.extract", + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ): task_metadata = self.extractor_manager.extract_metadata( dagrun=dagrun, @@ -284,6 +296,7 @@ def on_running(): "Skipping OpenLineage operator metadata extraction for task `%s` due to emission_policy.", task_instance.task_id, ) + task_metadata = OperatorLineage() redacted_event = self.adapter.start_task( @@ -318,10 +331,23 @@ def on_running(): }, ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) + + airflow_facet = redacted_event.run.facets.get("airflow") + team_name = None + + if airflow_facet: + team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) + Stats.gauge( "ol.event.size", event_size, - tags={"event_type": event_type, "operator_name": operator_name}, + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ) self._execute(on_running, "on_running", use_fork=True) @@ -413,9 +439,19 @@ def on_success(): if not doc: doc, doc_type = get_dag_documentation(dag) + team_name = None + team_name = DagRunInfo.team_name(dagrun) + if controls.extract_operator_metadata: with Stats.timer( - "ol.extract", tags={"event_type": event_type, "operator_name": operator_name} + "ol.extract", + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ): task_metadata = self.extractor_manager.extract_metadata( dagrun=dagrun, @@ -429,6 +465,7 @@ def on_success(): "Skipping OpenLineage operator metadata extraction for task `%s` due to emission_policy.", task_instance.task_id, ) + task_metadata = OperatorLineage() redacted_event = self.adapter.complete_task( @@ -462,10 +499,23 @@ def on_success(): }, ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) + + airflow_facet = redacted_event.run.facets.get("airflow") + team_name = None + + if airflow_facet: + team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) + Stats.gauge( "ol.event.size", event_size, - tags={"event_type": event_type, "operator_name": operator_name}, + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ) self._execute(on_success, "on_success", use_fork=True) @@ -572,9 +622,19 @@ def on_failure(): if not doc: doc, doc_type = get_dag_documentation(dag) + team_name = None + team_name = DagRunInfo.team_name(dagrun) + if controls.extract_operator_metadata: with Stats.timer( - "ol.extract", tags={"event_type": event_type, "operator_name": operator_name} + "ol.extract", + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ): task_metadata = self.extractor_manager.extract_metadata( dagrun=dagrun, @@ -622,10 +682,23 @@ def on_failure(): }, ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) + + airflow_facet = redacted_event.run.facets.get("airflow") + team_name = None + + if airflow_facet: + team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) + Stats.gauge( "ol.event.size", event_size, - tags={"event_type": event_type, "operator_name": operator_name}, + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ) self._execute(on_failure, "on_failure", use_fork=True) @@ -708,9 +781,19 @@ def on_skipped(): if not doc: doc, doc_type = get_dag_documentation(dag) + team_name = None + team_name = DagRunInfo.team_name(dagrun) + if controls.extract_operator_metadata: with Stats.timer( - "ol.extract", tags={"event_type": event_type, "operator_name": operator_name} + "ol.extract", + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ): task_metadata = self.extractor_manager.extract_metadata( dagrun=dagrun, @@ -757,10 +840,23 @@ def on_skipped(): }, ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) + + airflow_facet = redacted_event.run.facets.get("airflow") + team_name = None + + if airflow_facet: + team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) + Stats.gauge( "ol.event.size", event_size, - tags={"event_type": event_type, "operator_name": operator_name}, + tags=prune_dict( + { + "event_type": event_type, + "operator_name": operator_name, + "team_name": team_name, + } + ), ) self._execute(on_skipped, "on_skipped", use_fork=True) diff --git a/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py b/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py index 6bf7658754615..bbe08c38f2e9c 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py +++ b/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py @@ -24,7 +24,7 @@ from contextlib import suppress from functools import wraps from importlib import metadata -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, ClassVar import attrs from openlineage.client.facet_v2 import ( @@ -50,6 +50,7 @@ BaseOperator, BaseSensorOperator, MappedOperator, + conf as airflow_conf, ) from airflow.providers.openlineage import ( __version__ as OPENLINEAGE_PROVIDER_VERSION, @@ -72,6 +73,7 @@ from airflow.providers.openlineage.version_compat import ( AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_2_PLUS, + AIRFLOW_V_3_3_PLUS, get_base_airflow_version_tuple, ) from airflow.serialization.serialized_objects import SerializedBaseOperator, SerializedDAG @@ -86,6 +88,9 @@ if not AIRFLOW_V_3_0_PLUS: from airflow.utils.session import NEW_SESSION, provide_session +if AIRFLOW_V_3_3_PLUS: + from airflow.models.dagbundle import DagBundleModel + if TYPE_CHECKING: from typing import TypeAlias @@ -980,9 +985,12 @@ class DagRunInfo(InfoJsonEncodable): "dag_bundle_version": lambda dagrun: DagRunInfo.dag_version_info(dagrun, "bundle_version"), "dag_version_id": lambda dagrun: DagRunInfo.dag_version_info(dagrun, "version_id"), "dag_version_number": lambda dagrun: DagRunInfo.dag_version_info(dagrun, "version_number"), + "dag_team_name": lambda dagrun: DagRunInfo.team_name(dagrun) if AIRFLOW_V_3_3_PLUS else None, "deadlines": lambda dagrun: DagRunInfo.deadlines(dagrun), } + _team_name_cache: ClassVar[dict[str, str | None]] = {} + @classmethod def duration(cls, dagrun: DagRun) -> float | None: if not getattr(dagrun, "end_date", None) or not isinstance(dagrun.end_date, datetime.datetime): @@ -1053,6 +1061,21 @@ def dag_version_info(cls, dagrun: DagRun, key: str) -> str | int | None: return current_version.version_number raise ValueError(f"Unsupported key: {key}`") + @classmethod + def team_name(cls, dagrun: DagRun) -> str | None: + """Extract the team name for the DagRun.""" + if not AIRFLOW_V_3_3_PLUS or not airflow_conf.getboolean("core", "multi_team", fallback=False): + return None + + bundle_name = cls.dag_version_info(dagrun, "bundle_name") + if not isinstance(bundle_name, str): + return None + + if bundle_name not in cls._team_name_cache: + cls._team_name_cache[bundle_name] = DagBundleModel.get_team_name(bundle_name) + + return cls._team_name_cache[bundle_name] + class TaskInstanceInfo(InfoJsonEncodable): """Defines encoding TaskInstance object to JSON.""" diff --git a/providers/openlineage/src/airflow/providers/openlineage/version_compat.py b/providers/openlineage/src/airflow/providers/openlineage/version_compat.py index 114631640bc3c..663f36ccc16e0 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/version_compat.py +++ b/providers/openlineage/src/airflow/providers/openlineage/version_compat.py @@ -34,6 +34,7 @@ def get_base_airflow_version_tuple() -> tuple[int, int, int]: AIRFLOW_V_3_0_PLUS = get_base_airflow_version_tuple() >= (3, 0, 0) AIRFLOW_V_3_2_PLUS = get_base_airflow_version_tuple() >= (3, 2, 0) +AIRFLOW_V_3_3_PLUS = get_base_airflow_version_tuple() >= (3, 3, 0) -__all__ = ["AIRFLOW_V_3_0_PLUS", "AIRFLOW_V_3_2_PLUS"] +__all__ = ["AIRFLOW_V_3_0_PLUS", "AIRFLOW_V_3_2_PLUS", "AIRFLOW_V_3_3_PLUS"] diff --git a/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py b/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py index a3a252b8cbb18..5c8d32f0bc761 100644 --- a/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py +++ b/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py @@ -21,6 +21,7 @@ import os import pathlib import uuid +from types import SimpleNamespace from unittest import mock from unittest.mock import ANY, MagicMock, call, patch @@ -49,6 +50,7 @@ from airflow.providers.openlineage.plugins.facets import ( AirflowDagRunFacet, AirflowDebugRunFacet, + AirflowRunFacet, AirflowStateRunFacet, ) from airflow.providers.openlineage.token_provider import ( @@ -65,7 +67,7 @@ from tests_common.test_utils.config import conf_vars from tests_common.test_utils.markers import skip_if_force_lowest_dependencies_marker from tests_common.test_utils.taskinstance import create_task_instance -from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS +from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_3_PLUS stats_reference = f"{Stats.__module__}.Stats" @@ -304,14 +306,59 @@ def test_create_client_overrides_env_vars(): assert client.transport.kind == "console" +@pytest.mark.parametrize( + ("team_name", "expected_tags"), + [ + pytest.param( + None, + { + "event_type": "start", + "transport_type": ANY, + }, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": "start", + "transport_type": ANY, + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], +) @mock.patch(f"{stats_reference}.timer") @mock.patch(f"{stats_reference}.incr") -def test_emit_start_event(mock_stats_incr, mock_stats_timer): +def test_emit_start_event( + mock_stats_incr, + mock_stats_timer, + team_name, + expected_tags, +): + client = MagicMock() adapter = OpenLineageAdapter(client) run_id = str(uuid.uuid4()) event_time = datetime.datetime.now().isoformat() + + run_facets = None + if team_name is not None: + run_facets = { + "airflow": AirflowRunFacet( + dag={}, + dagRun={"dag_team_name": team_name}, + taskInstance={}, + task={}, + taskUuid="task_uuid", + ) + } + adapter.start_task( run_id=run_id, job_name="job", @@ -322,7 +369,7 @@ def test_emit_start_event(mock_stats_incr, mock_stats_timer): owners=[], tags=[], task=None, - run_facets=None, + run_facets=run_facets, ) assert ( @@ -340,6 +387,7 @@ def test_emit_start_event(mock_stats_incr, mock_stats_timer): "processing_engine": processing_engine_run.ProcessingEngineRunFacet( version=ANY, name="Airflow", openlineageAdapterVersion=ANY ), + **({"airflow": ANY} if team_name is not None else {}), }, ), job=Job( @@ -361,7 +409,10 @@ def test_emit_start_event(mock_stats_incr, mock_stats_timer): ) mock_stats_incr.assert_not_called() - mock_stats_timer.assert_called_with("ol.emit.attempts", tags={"event_type": ANY, "transport_type": ANY}) + mock_stats_timer.assert_called_with( + "ol.emit.attempts", + tags=expected_tags, + ) @mock.patch(f"{stats_reference}.timer") @@ -477,19 +528,64 @@ def test_emit_start_event_with_additional_information(mock_stats_incr, mock_stat mock_stats_timer.assert_called_with("ol.emit.attempts", tags={"event_type": ANY, "transport_type": ANY}) +@pytest.mark.parametrize( + ("team_name", "expected_tags"), + [ + pytest.param( + None, + { + "event_type": "complete", + "transport_type": ANY, + }, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": "complete", + "transport_type": ANY, + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], +) @mock.patch(f"{stats_reference}.timer") @mock.patch(f"{stats_reference}.incr") -def test_emit_complete_event(mock_stats_incr, mock_stats_timer): +def test_emit_complete_event( + mock_stats_incr, + mock_stats_timer, + team_name, + expected_tags, +): client = MagicMock() adapter = OpenLineageAdapter(client) run_id = str(uuid.uuid4()) event_time = datetime.datetime.now().isoformat() + + task = OperatorLineage() + + if team_name is not None: + task.run_facets = { + "airflow": AirflowRunFacet( + dag=None, + dagRun={"dag_team_name": team_name}, + taskInstance=None, + task=None, + taskUuid="task_uuid", + ) + } + adapter.complete_task( run_id=run_id, end_time=event_time, job_name="job", - task=OperatorLineage(), + task=task, owners=[], tags=[], job_description=None, @@ -505,8 +601,11 @@ def test_emit_complete_event(mock_stats_incr, mock_stats_timer): run=Run( runId=run_id, facets={ + **({"airflow": ANY} if team_name is not None else {}), "processing_engine": processing_engine_run.ProcessingEngineRunFacet( - version=ANY, name="Airflow", openlineageAdapterVersion=ANY + version=ANY, + name="Airflow", + openlineageAdapterVersion=ANY, ), "nominalTime": nominal_time_run.NominalTimeRunFacet( nominalStartTime="2022-01-01T00:00:00", @@ -519,7 +618,9 @@ def test_emit_complete_event(mock_stats_incr, mock_stats_timer): name="job", facets={ "jobType": job_type_job.JobTypeJobFacet( - processingType="BATCH", integration="AIRFLOW", jobType="TASK" + processingType="BATCH", + integration="AIRFLOW", + jobType="TASK", ) }, ), @@ -532,7 +633,10 @@ def test_emit_complete_event(mock_stats_incr, mock_stats_timer): ) mock_stats_incr.assert_not_called() - mock_stats_timer.assert_called_with("ol.emit.attempts", tags={"event_type": ANY, "transport_type": ANY}) + mock_stats_timer.assert_called_with( + "ol.emit.attempts", + tags=expected_tags, + ) @mock.patch(f"{stats_reference}.timer") @@ -650,19 +754,63 @@ def test_emit_complete_event_with_additional_information(mock_stats_incr, mock_s mock_stats_timer.assert_called_with("ol.emit.attempts", tags={"event_type": ANY, "transport_type": ANY}) +@pytest.mark.parametrize( + ("team_name", "expected_tags"), + [ + pytest.param( + None, + { + "event_type": "fail", + "transport_type": ANY, + }, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": "fail", + "transport_type": ANY, + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], +) @mock.patch(f"{stats_reference}.timer") @mock.patch(f"{stats_reference}.incr") -def test_emit_failed_event(mock_stats_incr, mock_stats_timer): +def test_emit_failed_event( + mock_stats_incr, + mock_stats_timer, + team_name, + expected_tags, +): client = MagicMock() adapter = OpenLineageAdapter(client) run_id = str(uuid.uuid4()) event_time = datetime.datetime.now().isoformat() + + task = OperatorLineage() + if team_name is not None: + task.run_facets = { + "airflow": AirflowRunFacet( + dag=None, + dagRun={"dag_team_name": team_name}, + taskInstance=None, + task=None, + taskUuid="task_uuid", + ) + } + adapter.fail_task( run_id=run_id, end_time=event_time, job_name="job", - task=OperatorLineage(), + task=task, owners=[], tags=[], job_description=None, @@ -678,6 +826,7 @@ def test_emit_failed_event(mock_stats_incr, mock_stats_timer): run=Run( runId=run_id, facets={ + **({"airflow": ANY} if team_name is not None else {}), "processing_engine": processing_engine_run.ProcessingEngineRunFacet( version=ANY, name="Airflow", openlineageAdapterVersion=ANY ), @@ -705,7 +854,10 @@ def test_emit_failed_event(mock_stats_incr, mock_stats_timer): ) mock_stats_incr.assert_not_called() - mock_stats_timer.assert_called_with("ol.emit.attempts", tags={"event_type": ANY, "transport_type": ANY}) + mock_stats_timer.assert_called_with( + "ol.emit.attempts", + tags=expected_tags, + ) @mock.patch(f"{stats_reference}.timer") @@ -1335,20 +1487,76 @@ def test_emit_dag_failed_event( mock_stats_timer.assert_called_with("ol.emit.attempts", tags={"event_type": ANY, "transport_type": ANY}) +@pytest.mark.parametrize( + ("team_name", "expected_timer_tags", "expected_failed_tags"), + [ + pytest.param( + None, + { + "event_type": ANY, + "transport_type": ANY, + }, + {}, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": ANY, + "transport_type": ANY, + "team_name": "team_a", + }, + { + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], +) @patch("airflow.providers.openlineage.plugins.adapter.OpenLineageAdapter.get_or_create_openlineage_client") @patch("airflow.providers.openlineage.plugins.adapter.OpenLineageRedactor") @patch(f"{stats_reference}.timer") @patch(f"{stats_reference}.incr") def test_openlineage_adapter_stats_emit_failed( - mock_stats_incr, mock_stats_timer, mock_redact, mock_get_client + mock_stats_incr, + mock_stats_timer, + mock_redact, + mock_get_client, + team_name, + expected_timer_tags, + expected_failed_tags, ): adapter = OpenLineageAdapter() mock_get_client.return_value.emit.side_effect = Exception() - adapter.emit(MagicMock()) + event = SimpleNamespace( + eventType=SimpleNamespace(value="COMPLETE"), + run=SimpleNamespace( + runId="run-id", + facets={}, + ), + ) - mock_stats_timer.assert_called_with("ol.emit.attempts", tags={"event_type": ANY, "transport_type": ANY}) - mock_stats_incr.assert_has_calls([mock.call("ol.emit.failed")]) + if team_name is not None: + event.run.facets["airflow"] = SimpleNamespace( + dagRun={"dag_team_name": team_name}, + ) + + adapter.emit(event) + + mock_stats_timer.assert_called_once_with( + "ol.emit.attempts", + tags=expected_timer_tags, + ) + + mock_stats_incr.assert_called_once_with( + "ol.emit.failed", + tags=expected_failed_tags, + ) def test_build_dag_run_id_is_valid_uuid(): diff --git a/providers/openlineage/tests/unit/openlineage/plugins/test_listener.py b/providers/openlineage/tests/unit/openlineage/plugins/test_listener.py index f7f3b398d4bfc..a4a31884a3ce2 100644 --- a/providers/openlineage/tests/unit/openlineage/plugins/test_listener.py +++ b/providers/openlineage/tests/unit/openlineage/plugins/test_listener.py @@ -23,6 +23,7 @@ from concurrent.futures import Future from contextlib import suppress from datetime import datetime +from types import SimpleNamespace from typing import TYPE_CHECKING from unittest import mock from unittest.mock import MagicMock, patch @@ -49,7 +50,12 @@ from tests_common.test_utils.dag import create_scheduler_dag from tests_common.test_utils.db import clear_db_runs from tests_common.test_utils.taskinstance import create_task_instance -from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_1_PLUS, AIRFLOW_V_3_2_PLUS +from tests_common.test_utils.version_compat import ( + AIRFLOW_V_3_0_PLUS, + AIRFLOW_V_3_1_PLUS, + AIRFLOW_V_3_2_PLUS, + AIRFLOW_V_3_3_PLUS, +) if AIRFLOW_V_3_1_PLUS: from airflow._shared.timezones import timezone @@ -1380,6 +1386,35 @@ def mock_task_id(dag_id, task_id, try_number, logical_date, map_index): return listener, task_instance + @pytest.mark.parametrize( + ("team_name", "expected_tags"), + [ + pytest.param( + None, + { + "event_type": "running", + "operator_name": "emptyoperator", + }, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": "running", + "operator_name": "emptyoperator", + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], + ) + @mock.patch("airflow.providers.openlineage.plugins.listener.Serde.to_json") + @mock.patch("airflow.providers.openlineage.plugins.listener.DagRunInfo.team_name") + @mock.patch("airflow.providers.openlineage.plugins.listener.Stats") @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_debug_facet") @mock.patch("airflow.providers.openlineage.plugins.listener.resolve_task_emission_policy") @@ -1399,6 +1434,11 @@ def test_adapter_start_task_is_called_with_proper_arguments( mock_disabled, mock_debug_facet, mock_debug_mode, + mock_stats, + mock_team_name, + mock_to_json, + team_name, + expected_tags, ): """Tests that the 'start_task' method of the OpenLineageAdapter is invoked with the correct arguments. @@ -1416,6 +1456,26 @@ def test_adapter_start_task_is_called_with_proper_arguments( mock_get_task_parent_run_facet.return_value = {"parent": 4} mock_debug_facet.return_value = {"debug": "packages"} mock_disabled.return_value = EmissionPolicy.defaults() + mock_team_name.return_value = team_name + mock_to_json.return_value = "{}" + + fake_event = SimpleNamespace( + run=SimpleNamespace( + facets=( + {} + if team_name is None + else { + "airflow": SimpleNamespace( + dagRun=SimpleNamespace( + dag_team_name=team_name, + ) + ) + } + ) + ) + ) + + listener.adapter.start_task.return_value = fake_event listener.on_task_instance_running(None, task_instance) listener.adapter.start_task.assert_called_once_with( @@ -1438,6 +1498,17 @@ def test_adapter_start_task_is_called_with_proper_arguments( }, ) + mock_stats.timer.assert_any_call( + "ol.extract", + tags=expected_tags, + ) + + mock_stats.gauge.assert_called_once_with( + "ol.event.size", + mock.ANY, + tags=expected_tags, + ) + @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_debug_facet") @mock.patch("airflow.providers.openlineage.plugins.listener.get_task_parent_run_facet") @@ -1680,6 +1751,35 @@ def test_adapter_start_task_is_called_with_dag_description_when_task_doc_is_empt assert listener.adapter.start_task.call_args.kwargs["job_description"] == "Test DAG Description" assert listener.adapter.start_task.call_args.kwargs["job_description_type"] == "text/plain" + @pytest.mark.parametrize( + ("team_name", "expected_tags"), + [ + pytest.param( + None, + { + "event_type": "fail", + "operator_name": "emptyoperator", + }, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": "fail", + "operator_name": "emptyoperator", + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], + ) + @mock.patch("airflow.providers.openlineage.plugins.listener.Serde.to_json") + @mock.patch("airflow.providers.openlineage.plugins.listener.DagRunInfo.team_name") + @mock.patch("airflow.providers.openlineage.plugins.listener.Stats") @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_debug_facet") @mock.patch("airflow.providers.openlineage.plugins.listener.resolve_task_emission_policy") @@ -1697,7 +1797,12 @@ def test_adapter_fail_task_is_called_with_proper_arguments( mock_disabled, mock_debug_facet, mock_debug_mode, + mock_stats, + mock_team_name, + mock_to_json, time_machine, + team_name, + expected_tags, ): """Tests that the 'fail_task' method of the OpenLineageAdapter is invoked with the correct arguments. @@ -1715,8 +1820,29 @@ def test_adapter_fail_task_is_called_with_proper_arguments( mock_get_task_parent_run_facet.return_value = {"parent": 4} mock_debug_facet.return_value = {"debug": "packages"} mock_disabled.return_value = EmissionPolicy.defaults() + mock_team_name.return_value = team_name + mock_to_json.return_value = "{}" err = ValueError("test") + + fake_event = SimpleNamespace( + run=SimpleNamespace( + facets=( + {} + if team_name is None + else { + "airflow": SimpleNamespace( + dagRun=SimpleNamespace( + dag_team_name=team_name, + ) + ) + } + ) + ) + ) + + listener.adapter.fail_task.return_value = fake_event + listener.on_task_instance_failed(previous_state=None, task_instance=task_instance, error=err) listener.adapter.fail_task.assert_called_once_with( end_time="2023-01-03T13:01:01+00:00", @@ -1738,6 +1864,17 @@ def test_adapter_fail_task_is_called_with_proper_arguments( error=err, ) + mock_stats.timer.assert_any_call( + "ol.extract", + tags=expected_tags, + ) + + mock_stats.gauge.assert_called_once_with( + "ol.event.size", + mock.ANY, + tags=expected_tags, + ) + @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_debug_facet") @mock.patch("airflow.providers.openlineage.plugins.listener.resolve_task_emission_policy") @@ -1875,6 +2012,35 @@ def test_adapter_fail_task_is_called_with_proper_arguments_for_db_task_instance_ adapter.fail_task(**expected_args) assert mock_emit.assert_called_once + @pytest.mark.parametrize( + ("team_name", "expected_tags"), + [ + pytest.param( + None, + { + "event_type": "complete", + "operator_name": "emptyoperator", + }, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": "complete", + "operator_name": "emptyoperator", + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], + ) + @mock.patch("airflow.providers.openlineage.plugins.listener.Serde.to_json") + @mock.patch("airflow.providers.openlineage.plugins.listener.DagRunInfo.team_name") + @mock.patch("airflow.providers.openlineage.plugins.listener.Stats") @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_debug_facet") @mock.patch("airflow.providers.openlineage.plugins.listener.resolve_task_emission_policy") @@ -1892,7 +2058,12 @@ def test_adapter_complete_task_is_called_with_proper_arguments( mock_disabled, mock_debug_facet, mock_debug_mode, + mock_stats, + mock_team_name, + mock_to_json, time_machine, + team_name, + expected_tags, ): """Tests that the 'complete_task' method of the OpenLineageAdapter is called with the correct arguments. @@ -1910,6 +2081,26 @@ def test_adapter_complete_task_is_called_with_proper_arguments( mock_get_task_parent_run_facet.return_value = {"parent": 4} mock_debug_facet.return_value = {"debug": "packages"} mock_disabled.return_value = EmissionPolicy.defaults() + mock_team_name.return_value = team_name + mock_to_json.return_value = "{}" + + fake_event = SimpleNamespace( + run=SimpleNamespace( + facets=( + {} + if team_name is None + else { + "airflow": SimpleNamespace( + dagRun=SimpleNamespace( + dag_team_name=team_name, + ) + ) + } + ) + ) + ) + + listener.adapter.complete_task.return_value = fake_event listener.on_task_instance_success(None, task_instance) calls = listener.adapter.complete_task.call_args_list @@ -1933,6 +2124,17 @@ def test_adapter_complete_task_is_called_with_proper_arguments( }, ) + mock_stats.timer.assert_any_call( + "ol.extract", + tags=expected_tags, + ) + + mock_stats.gauge.assert_called_once_with( + "ol.event.size", + mock.ANY, + tags=expected_tags, + ) + @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_debug_facet") @mock.patch("airflow.providers.openlineage.plugins.listener.resolve_task_emission_policy") @@ -2242,6 +2444,121 @@ def test_listener_on_task_instance_skipped_do_not_call_adapter_when_disabled_ope listener.extractor_manager.extract_metadata.assert_not_called() listener.adapter.complete_task.assert_not_called() + @pytest.mark.parametrize( + ("team_name", "expected_tags"), + [ + pytest.param( + None, + { + "event_type": "complete", + "operator_name": "emptyoperator", + }, + id="without_team", + ), + pytest.param( + "team_a", + { + "event_type": "complete", + "operator_name": "emptyoperator", + "team_name": "team_a", + }, + id="with_team", + marks=pytest.mark.skipif( + not AIRFLOW_V_3_3_PLUS, + reason="team_name metrics require Airflow 3.3+", + ), + ), + ], + ) + @mock.patch("airflow.providers.openlineage.plugins.listener.Serde.to_json") + @mock.patch("airflow.providers.openlineage.plugins.listener.DagRunInfo.team_name") + @mock.patch("airflow.providers.openlineage.plugins.listener.Stats") + @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) + @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_debug_facet") + @mock.patch("airflow.providers.openlineage.plugins.listener.resolve_task_emission_policy") + @mock.patch("airflow.providers.openlineage.plugins.listener.get_task_parent_run_facet") + @mock.patch("airflow.providers.openlineage.plugins.listener.get_airflow_run_facet") + @mock.patch("airflow.providers.openlineage.plugins.listener.get_user_provided_run_facets") + @mock.patch( + "airflow.providers.openlineage.plugins.listener.OpenLineageListener._execute", + new=regular_call, + ) + def test_adapter_complete_task_is_called_with_proper_arguments_on_skip( + self, + mock_get_user_provided_run_facets, + mock_get_airflow_run_facet, + mock_get_task_parent_run_facet, + mock_disabled, + mock_debug_facet, + mock_debug_mode, + mock_stats, + mock_team_name, + mock_to_json, + time_machine, + team_name, + expected_tags, + ): + time_machine.move_to(timezone.datetime(2023, 1, 3, 13, 1, 1), tick=False) + + listener, task_instance = self._create_listener_and_task_instance() + + mock_get_user_provided_run_facets.return_value = {"custom_user_facet": 2, "parent": 99} + mock_get_airflow_run_facet.return_value = {"airflow": {"task": "..."}} + mock_get_task_parent_run_facet.return_value = {"parent": 4} + mock_debug_facet.return_value = {"debug": "packages"} + mock_disabled.return_value = EmissionPolicy.defaults() + mock_team_name.return_value = team_name + mock_to_json.return_value = "{}" + + fake_event = SimpleNamespace( + run=SimpleNamespace( + facets=( + {} + if team_name is None + else { + "airflow": SimpleNamespace( + dagRun=SimpleNamespace( + dag_team_name=team_name, + ) + ) + } + ) + ) + ) + listener.adapter.complete_task.return_value = fake_event + + listener.on_task_instance_skipped(previous_state=None, task_instance=task_instance) + + listener.adapter.complete_task.assert_called_once_with( + end_time="2023-01-03T13:01:01+00:00", + job_name="dag_id.task_id", + run_id="2020-01-01T01:01:01+00:00.dag_id.task_id.1.-1", + task=listener.extractor_manager.extract_metadata(), + owners=["task_owner"], + tags={"tag1", "tag2"}, + job_description="TASK Description", + job_description_type="text/markdown", + nominal_start_time=None, + nominal_end_time=None, + run_facets={ + "parent": 4, + "custom_user_facet": 2, + "airflow": {"task": "..."}, + "debug": "packages", + }, + ) + + mock_stats.timer.assert_any_call( + "ol.extract", + tags=expected_tags, + ) + + mock_stats.gauge.assert_called_once_with( + "ol.event.size", + mock.ANY, + tags=expected_tags, + ) + @mock.patch("airflow.providers.openlineage.plugins.listener.OpenLineageListener._fork_execute") @mock.patch("airflow.providers.openlineage.plugins.adapter.OpenLineageAdapter.emit") @mock.patch("airflow.providers.openlineage.conf.debug_mode", return_value=True) diff --git a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py index 36fe5125a7f0f..767690b1523aa 100644 --- a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py +++ b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py @@ -92,6 +92,7 @@ AIRFLOW_V_3_0_3_PLUS, AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_2_PLUS, + AIRFLOW_V_3_3_PLUS, ) BASH_OPERATOR_PATH = "airflow.providers.standard.operators.bash" @@ -276,6 +277,7 @@ def test_get_airflow_dag_run_facet(): "dag_bundle_version": "bundle_version", "dag_version_id": "version_id", "dag_version_number": "version_number", + "dag_team_name": None, "triggering_user_name": "user1", "partition_key": "some_partition_key", "partition_date": "2024-06-01T02:03:34+00:00", @@ -331,6 +333,57 @@ def test_dag_run_version(key): assert result == key +@pytest.mark.db_test +@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") +@patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") +@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True) +def test_dag_run_team_name( + mock_getboolean, + mock_get_team_name, +): + DagRunInfo._team_name_cache.clear() + + dagrun_mock = MagicMock(DagRun) + dagrun_mock.dag_versions = [ + MagicMock( + bundle_name="bundle_name", + bundle_version="bundle_version", + id="version_id", + version_number="version_number", + ) + ] + + mock_get_team_name.return_value = "team_a" + + assert DagRunInfo.team_name(dagrun_mock) == "team_a" + assert DagRunInfo.team_name(dagrun_mock) == "team_a" + + mock_get_team_name.assert_called_once_with("bundle_name") + + +@pytest.mark.db_test +@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.1+") +def test_dag_run_team_name_no_bundle(): + dagrun_mock = MagicMock(DagRun) + del dagrun_mock.dag_versions + + assert DagRunInfo.team_name(dagrun_mock) is None + + +@pytest.mark.db_test +@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.1+") +@patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") +@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=False) +def test_dag_run_team_name_multi_team_disabled(mock_getboolean, mock_get_team_name): + + DagRunInfo._team_name_cache.clear() + + dagrun_mock = MagicMock(DagRun) + + assert DagRunInfo.team_name(dagrun_mock) is None + mock_get_team_name.assert_not_called() + + def test_get_fully_qualified_class_name_serialized_operator(): op_module_path = BASH_OPERATOR_PATH op_name = "BashOperator" @@ -2965,6 +3018,7 @@ def test_dagrun_info_af3(mocked_dag_versions): "conf": {"a": 1}, "clear_number": 0, "dag_id": "dag_id", + "dag_team_name": None, "data_interval_end": "2024-06-01T00:00:00+00:00", "data_interval_start": "2024-06-01T00:00:00+00:00", "duration": 74.000546, @@ -3011,6 +3065,7 @@ def test_dagrun_info_af2(): "conf": {"a": 1}, "clear_number": 0, "dag_id": "dag_id", + "dag_team_name": None, "data_interval_end": "2024-06-01T00:00:00+00:00", "data_interval_start": "2024-06-01T00:00:00+00:00", "duration": 74.000546, From 59a92b8ddf1ff88c96496b44c65377984787e3cc Mon Sep 17 00:00:00 2001 From: Sameer Mesiah Date: Mon, 6 Jul 2026 22:42:31 +0100 Subject: [PATCH 2/5] Correct the Airflow version referenced in the skipif reason strings and reuse the existing team_name value when emitting metrics. --- .../providers/openlineage/plugins/listener.py | 28 ------------------- .../unit/openlineage/utils/test_utils.py | 4 +-- 2 files changed, 2 insertions(+), 30 deletions(-) diff --git a/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py b/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py index 38f0f7ad060e2..6ce03e6f07862 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py +++ b/providers/openlineage/src/airflow/providers/openlineage/plugins/listener.py @@ -270,7 +270,6 @@ def on_running(): if not doc: doc, doc_type = get_dag_documentation(dag) - team_name = None team_name = DagRunInfo.team_name(dagrun) if controls.extract_operator_metadata: @@ -332,12 +331,6 @@ def on_running(): ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) - airflow_facet = redacted_event.run.facets.get("airflow") - team_name = None - - if airflow_facet: - team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) - Stats.gauge( "ol.event.size", event_size, @@ -439,7 +432,6 @@ def on_success(): if not doc: doc, doc_type = get_dag_documentation(dag) - team_name = None team_name = DagRunInfo.team_name(dagrun) if controls.extract_operator_metadata: @@ -500,12 +492,6 @@ def on_success(): ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) - airflow_facet = redacted_event.run.facets.get("airflow") - team_name = None - - if airflow_facet: - team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) - Stats.gauge( "ol.event.size", event_size, @@ -622,7 +608,6 @@ def on_failure(): if not doc: doc, doc_type = get_dag_documentation(dag) - team_name = None team_name = DagRunInfo.team_name(dagrun) if controls.extract_operator_metadata: @@ -683,12 +668,6 @@ def on_failure(): ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) - airflow_facet = redacted_event.run.facets.get("airflow") - team_name = None - - if airflow_facet: - team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) - Stats.gauge( "ol.event.size", event_size, @@ -781,7 +760,6 @@ def on_skipped(): if not doc: doc, doc_type = get_dag_documentation(dag) - team_name = None team_name = DagRunInfo.team_name(dagrun) if controls.extract_operator_metadata: @@ -841,12 +819,6 @@ def on_skipped(): ) event_size = len(Serde.to_json(redacted_event).encode("utf-8")) - airflow_facet = redacted_event.run.facets.get("airflow") - team_name = None - - if airflow_facet: - team_name = getattr(airflow_facet.dagRun, "dag_team_name", None) - Stats.gauge( "ol.event.size", event_size, diff --git a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py index 767690b1523aa..feb392a76d750 100644 --- a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py +++ b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py @@ -362,7 +362,7 @@ def test_dag_run_team_name( @pytest.mark.db_test -@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.1+") +@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") def test_dag_run_team_name_no_bundle(): dagrun_mock = MagicMock(DagRun) del dagrun_mock.dag_versions @@ -371,7 +371,7 @@ def test_dag_run_team_name_no_bundle(): @pytest.mark.db_test -@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.1+") +@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") @patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") @patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=False) def test_dag_run_team_name_multi_team_disabled(mock_getboolean, mock_get_team_name): From 8c4cd88d8da376fb253145f2231ba5f68e776491 Mon Sep 17 00:00:00 2001 From: Sameer Mesiah Date: Wed, 8 Jul 2026 21:02:16 +0100 Subject: [PATCH 3/5] Remove the in-memory cache from DagRunInfo.team_name and refine the associated unit tests by removing unnecessary database markers and simplifying the test cases. --- .../airflow/providers/openlineage/utils/utils.py | 11 +++-------- .../tests/unit/openlineage/utils/test_utils.py | 13 ++++++------- 2 files changed, 9 insertions(+), 15 deletions(-) diff --git a/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py b/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py index bbe08c38f2e9c..5932cab0854df 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py +++ b/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py @@ -24,7 +24,7 @@ from contextlib import suppress from functools import wraps from importlib import metadata -from typing import TYPE_CHECKING, Any, ClassVar +from typing import TYPE_CHECKING, Any import attrs from openlineage.client.facet_v2 import ( @@ -989,8 +989,6 @@ class DagRunInfo(InfoJsonEncodable): "deadlines": lambda dagrun: DagRunInfo.deadlines(dagrun), } - _team_name_cache: ClassVar[dict[str, str | None]] = {} - @classmethod def duration(cls, dagrun: DagRun) -> float | None: if not getattr(dagrun, "end_date", None) or not isinstance(dagrun.end_date, datetime.datetime): @@ -1045,7 +1043,7 @@ def deadlines(cls, dagrun: DagRun) -> dict[str, Any] | None: @classmethod def dag_version_info(cls, dagrun: DagRun, key: str) -> str | int | None: - """Extract deg version info for given key, sourced from DagRun (on scheduler).""" + """Extract DAG version info for given key, sourced from DagRun (on scheduler).""" # AF2 DagRun and AF3 DagRun SDK model (on worker) do not have this information dag_versions = safe_getattr(dagrun, "dag_versions", []) if not dag_versions: @@ -1071,10 +1069,7 @@ def team_name(cls, dagrun: DagRun) -> str | None: if not isinstance(bundle_name, str): return None - if bundle_name not in cls._team_name_cache: - cls._team_name_cache[bundle_name] = DagBundleModel.get_team_name(bundle_name) - - return cls._team_name_cache[bundle_name] + return DagBundleModel.get_team_name(bundle_name) class TaskInstanceInfo(InfoJsonEncodable): diff --git a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py index feb392a76d750..691237545060b 100644 --- a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py +++ b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py @@ -341,7 +341,6 @@ def test_dag_run_team_name( mock_getboolean, mock_get_team_name, ): - DagRunInfo._team_name_cache.clear() dagrun_mock = MagicMock(DagRun) dagrun_mock.dag_versions = [ @@ -355,32 +354,32 @@ def test_dag_run_team_name( mock_get_team_name.return_value = "team_a" - assert DagRunInfo.team_name(dagrun_mock) == "team_a" assert DagRunInfo.team_name(dagrun_mock) == "team_a" mock_get_team_name.assert_called_once_with("bundle_name") -@pytest.mark.db_test @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") -def test_dag_run_team_name_no_bundle(): +@patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") +@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True) +def test_dag_run_team_name_no_bundle(mock_getboolean, mock_get_team_name): dagrun_mock = MagicMock(DagRun) del dagrun_mock.dag_versions assert DagRunInfo.team_name(dagrun_mock) is None + mock_get_team_name.assert_not_called() + -@pytest.mark.db_test @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") @patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") @patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=False) def test_dag_run_team_name_multi_team_disabled(mock_getboolean, mock_get_team_name): - DagRunInfo._team_name_cache.clear() - dagrun_mock = MagicMock(DagRun) assert DagRunInfo.team_name(dagrun_mock) is None + mock_get_team_name.assert_not_called() From e7835f6714a943b6417507e177d520b77b68bff6 Mon Sep 17 00:00:00 2001 From: Sameer Mesiah Date: Thu, 9 Jul 2026 21:54:25 +0100 Subject: [PATCH 4/5] Handle team_name extraction from airflowDagRun facets, add coverage for the fallback path, and restore the required db_test markers for the DagRunInfo tests. --- .../providers/openlineage/plugins/adapter.py | 11 ++++- .../unit/openlineage/plugins/test_adapter.py | 44 +++++++++++++++++++ .../unit/openlineage/utils/test_utils.py | 2 + 3 files changed, 56 insertions(+), 1 deletion(-) diff --git a/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py b/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py index 617395a989d06..b4bf365547bab 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py +++ b/providers/openlineage/src/airflow/providers/openlineage/plugins/adapter.py @@ -57,7 +57,7 @@ from datetime import datetime from airflow.providers.openlineage.extractors import OperatorLineage - from airflow.providers.openlineage.plugins.facets import AirflowRunFacet + from airflow.providers.openlineage.plugins.facets import AirflowDagRunFacet, AirflowRunFacet from airflow.sdk.execution_time.secrets_masker import SecretsMasker, _secrets_masker from airflow.utils.state import DagRunState else: @@ -181,6 +181,15 @@ def emit(self, event: RunEvent): if airflow_facet: team_name = airflow_facet.dagRun.get("dag_team_name") + else: + airflow_dagrun_facet = cast("AirflowDagRunFacet | None", facets.get("airflowDagRun")) + if airflow_dagrun_facet: + dag_run = airflow_dagrun_facet.dagRun + team_name = ( + dag_run.get("dag_team_name") + if isinstance(dag_run, dict) + else getattr(dag_run, "dag_team_name", None) + ) try: with Stats.timer( diff --git a/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py b/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py index 5c8d32f0bc761..e149188d7ef49 100644 --- a/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py +++ b/providers/openlineage/tests/unit/openlineage/plugins/test_adapter.py @@ -639,6 +639,50 @@ def test_emit_complete_event( ) +@mock.patch(f"{stats_reference}.timer") +@mock.patch(f"{stats_reference}.incr") +def test_emit_complete_event_dagrun_fallback( + mock_stats_incr, + mock_stats_timer, +): + client = MagicMock() + adapter = OpenLineageAdapter(client) + + event = RunEvent( + eventType=RunState.COMPLETE, + eventTime=datetime.datetime.now().isoformat(), + run=Run( + runId=str(uuid.uuid4()), + facets={ + "airflowDagRun": AirflowDagRunFacet( + dag={}, + dagRun={"dag_team_name": "team_a"}, + ), + }, + ), + job=Job( + namespace=namespace(), + name="dag", + facets={}, + ), + producer=_PRODUCER, + inputs=[], + outputs=[], + ) + + adapter.emit(event) + + mock_stats_incr.assert_not_called() + mock_stats_timer.assert_called_once_with( + "ol.emit.attempts", + tags={ + "event_type": "complete", + "transport_type": ANY, + "team_name": "team_a", + }, + ) + + @mock.patch(f"{stats_reference}.timer") @mock.patch(f"{stats_reference}.incr") def test_emit_complete_event_with_additional_information(mock_stats_incr, mock_stats_timer): diff --git a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py index 691237545060b..1599770af759e 100644 --- a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py +++ b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py @@ -359,6 +359,7 @@ def test_dag_run_team_name( mock_get_team_name.assert_called_once_with("bundle_name") +@pytest.mark.db_test @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") @patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") @patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True) @@ -371,6 +372,7 @@ def test_dag_run_team_name_no_bundle(mock_getboolean, mock_get_team_name): mock_get_team_name.assert_not_called() +@pytest.mark.db_test @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") @patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") @patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=False) From c338d5bddceb4f9014a8652ed955054091dc324c Mon Sep 17 00:00:00 2001 From: Sameer Mesiah Date: Fri, 10 Jul 2026 21:40:02 +0100 Subject: [PATCH 5/5] Move DagBundleModel import into team_name helper --- .../src/airflow/providers/openlineage/utils/utils.py | 5 ++--- .../openlineage/tests/unit/openlineage/utils/test_utils.py | 6 +++--- 2 files changed, 5 insertions(+), 6 deletions(-) diff --git a/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py b/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py index 5932cab0854df..f40e9b12a0be7 100644 --- a/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py +++ b/providers/openlineage/src/airflow/providers/openlineage/utils/utils.py @@ -88,9 +88,6 @@ if not AIRFLOW_V_3_0_PLUS: from airflow.utils.session import NEW_SESSION, provide_session -if AIRFLOW_V_3_3_PLUS: - from airflow.models.dagbundle import DagBundleModel - if TYPE_CHECKING: from typing import TypeAlias @@ -1065,6 +1062,8 @@ def team_name(cls, dagrun: DagRun) -> str | None: if not AIRFLOW_V_3_3_PLUS or not airflow_conf.getboolean("core", "multi_team", fallback=False): return None + from airflow.models.dagbundle import DagBundleModel + bundle_name = cls.dag_version_info(dagrun, "bundle_name") if not isinstance(bundle_name, str): return None diff --git a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py index 1599770af759e..e659e4f3de3da 100644 --- a/providers/openlineage/tests/unit/openlineage/utils/test_utils.py +++ b/providers/openlineage/tests/unit/openlineage/utils/test_utils.py @@ -335,7 +335,7 @@ def test_dag_run_version(key): @pytest.mark.db_test @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") -@patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") +@patch("airflow.models.dagbundle.DagBundleModel.get_team_name") @patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True) def test_dag_run_team_name( mock_getboolean, @@ -361,7 +361,7 @@ def test_dag_run_team_name( @pytest.mark.db_test @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") -@patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") +@patch("airflow.models.dagbundle.DagBundleModel.get_team_name") @patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True) def test_dag_run_team_name_no_bundle(mock_getboolean, mock_get_team_name): dagrun_mock = MagicMock(DagRun) @@ -374,7 +374,7 @@ def test_dag_run_team_name_no_bundle(mock_getboolean, mock_get_team_name): @pytest.mark.db_test @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+") -@patch("airflow.providers.openlineage.utils.utils.DagBundleModel.get_team_name") +@patch("airflow.models.dagbundle.DagBundleModel.get_team_name") @patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=False) def test_dag_run_team_name_multi_team_disabled(mock_getboolean, mock_get_team_name):