Skip to content
Merged
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 @@ -982,6 +982,9 @@ 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_version_data": lambda dagrun: (
DagRunInfo.dag_version_info(dagrun, "version_data") if AIRFLOW_V_3_3_PLUS else None
),
"dag_team_name": lambda dagrun: DagRunInfo.team_name(dagrun) if AIRFLOW_V_3_3_PLUS else None,
"deadlines": lambda dagrun: DagRunInfo.deadlines(dagrun),
}
Expand Down Expand Up @@ -1039,7 +1042,7 @@ def deadlines(cls, dagrun: DagRun) -> dict[str, Any] | None:
return {"alerts": result} if result else None

@classmethod
def dag_version_info(cls, dagrun: DagRun, key: str) -> str | int | None:
def dag_version_info(cls, dagrun: DagRun, key: str) -> str | int | dict | None:
"""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", [])
Expand All @@ -1057,6 +1060,10 @@ def dag_version_info(cls, dagrun: DagRun, key: str) -> str | int | None:
return str(version_id) if version_id is not None else None
if key == "version_number":
return safe_getattr(current_version, "version_number")
if key == "version_data":
if not AIRFLOW_V_3_3_PLUS:
return None
return safe_getattr(current_version, "version_data")
raise ValueError(f"Unsupported key: {key}`")

@classmethod
Expand All @@ -1073,18 +1080,29 @@ def team_name(cls, dagrun: DagRun) -> str | None:
if hasattr(dagrun, "team_name"):
return dagrun.team_name

# Best-effort: the scheduler stamps `_team_name` on ORM DagRun objects before
# listener hooks fire. It's a private attribute with no stability guarantee,
# so guard with hasattr and an isinstance check.
if hasattr(dagrun, "_team_name"):
return dagrun._team_name if isinstance(dagrun._team_name, str) else None

try:
bundle_name = cls.dag_version_info(dagrun, "bundle_name")
if not isinstance(bundle_name, str):
# Reuse the existing ORM session associated with the DagRun. Creating a new session here
# (via @provide_session) on get_team_name() can trigger an unexpected commit error.
from sqlalchemy.orm import object_session

session = object_session(dagrun)
Comment thread
kacpermuda marked this conversation as resolved.

if session is None:
return None

from airflow.models.dagbundle import DagBundleModel
from airflow.models.dag import DagModel

return DagBundleModel.get_team_name(bundle_name)
return DagModel.get_team_name(dagrun.dag_id, session=session)
Comment thread
kacpermuda marked this conversation as resolved.
except Exception as e:
log.warning(
log.info(
"OpenLineage failed to resolve the team name for dag `%s`: %s.",
safe_getattr(dagrun, "dag_id"),
dagrun.dag_id,
e,
)
log.debug("Exception details:", exc_info=True)
Expand All @@ -1094,9 +1112,10 @@ def team_name(cls, dagrun: DagRun) -> str | None:
class TaskInstanceInfo(InfoJsonEncodable):
"""Defines encoding TaskInstance object to JSON."""

includes = ["duration", "try_number", "pool", "queued_dttm", "log_url"]
includes = ["duration", "log_url", "pool", "queued_dttm", "try_number"]
casts = {
"log_url": lambda ti: getattr(ti, "log_url", None),
"note": lambda ti: safe_getattr(ti, "note", None), # From manual state changes only
"map_index": lambda ti: ti.map_index if getattr(ti, "map_index", -1) != -1 else None,
"rendered_map_index": lambda ti: (
getattr(ti, "rendered_map_index", None) if getattr(ti, "map_index", -1) != -1 else None
Expand Down
154 changes: 113 additions & 41 deletions providers/openlineage/tests/unit/openlineage/utils/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -233,6 +233,7 @@ def test_get_airflow_dag_run_facet():
bundle_version="bundle_version",
id="version_id",
version_number="version_number",
version_data={"some": "data"},
)
]
dagrun_mock.deadlines = []
Expand All @@ -252,9 +253,14 @@ def test_get_airflow_dag_run_facet():
}
if hasattr(dag, "schedule_interval"): # Airflow 2 compat.
expected_dag_info["schedule_interval"] = "@once"
note: str | None = None

optional_result = {}
if AIRFLOW_V_3_2_PLUS:
note = "note"
optional_result["note"] = "note"

if AIRFLOW_V_3_3_PLUS:
optional_result["dag_version_data"] = {"some": "data"}

assert result == {
"airflowDagRun": AirflowDagRunFacet(
dag=expected_dag_info,
Expand Down Expand Up @@ -283,7 +289,9 @@ def test_get_airflow_dag_run_facet():
"partition_key": "some_partition_key",
"partition_date": "2024-06-01T02:03:34+00:00",
"triggered_by": "something",
"note": note,
"note": None,
"dag_version_data": None,
**optional_result,
},
)
}
Expand Down Expand Up @@ -348,74 +356,92 @@ 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.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,
mock_get_team_name,
):

@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="version_data requires Airflow 3.3+")
def test_dag_run_version_data():
dagrun_mock = MagicMock(DagRun)
dagrun_mock.dag_versions = [
MagicMock(
bundle_name="bundle_name",
bundle_version="bundle_version",
id="version_id",
version_number="version_number",
)
]
dagrun_mock.dag_versions = [MagicMock(version_data={"schema": 1})]
assert DagRunInfo.dag_version_info(dagrun_mock, "version_data") == {"schema": 1}

mock_get_team_name.return_value = "team_a"

assert DagRunInfo.team_name(dagrun_mock) == "team_a"
@pytest.mark.db_test
@patch("airflow.providers.openlineage.utils.utils.AIRFLOW_V_3_3_PLUS", False)
def test_dag_run_version_data_below_3_3():
dagrun_mock = MagicMock(DagRun)
dagrun_mock.dag_versions = [MagicMock(version_data={"schema": 1})]
assert DagRunInfo.dag_version_info(dagrun_mock, "version_data") is None

mock_get_team_name.assert_called_once_with("bundle_name")

@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="version_data requires Airflow 3.3+")
def test_dag_run_version_data_detached_version_row():
"""version_data is lazy-loaded and can hit a detached session like the other columns."""
version = MagicMock()
type(version).version_data = PropertyMock(side_effect=DetachedInstanceError)
dag_run = MagicMock()
dag_run.dag_versions = [version]
assert DagRunInfo.dag_version_info(dag_run, "version_data") is None


@pytest.mark.db_test
@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+")
@patch("airflow.models.dag.DagModel.get_team_name", return_value="team_a")
@patch("sqlalchemy.orm.object_session")
@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True)
def test_dag_run_team_name(mock_getboolean, mock_object_session, mock_get_team_name):
"""DB fallback uses the dagrun's existing session — no new session opened, no HA lock risk."""
mock_session = MagicMock()
mock_object_session.return_value = mock_session

dagrun_mock = MagicMock(spec=DagRun)
dagrun_mock.dag_id = "test_dag"

assert DagRunInfo.team_name(dagrun_mock) == "team_a"
mock_get_team_name.assert_called_once_with("test_dag", session=mock_session)


@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+")
@pytest.mark.parametrize("team_name", ["team_a", None])
@patch("airflow.models.dagbundle.DagBundleModel.get_team_name")
@patch("sqlalchemy.orm.object_session")
@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True)
def test_dag_run_team_name_from_execution_api_dag_run(mock_getboolean, mock_get_team_name, team_name):
"""The task runner has no DB session, so a DagRun carrying `team_name` must be trusted as-is."""
# A resolvable bundle plus a DB answer, so falling through to the lookup would be observable.
dagrun_mock = MagicMock(spec_set=["team_name", "dag_versions"])
def test_dag_run_team_name_from_execution_api_dag_run(mock_getboolean, mock_object_session, team_name):
"""The task runner has no DB session, so a DagRun carrying `team_name` must be trusted as-is.

The cascade must stop at the first step — the DB lookup (object_session) must never be reached.
"""
dagrun_mock = MagicMock(spec_set=["team_name"])
dagrun_mock.team_name = team_name
dagrun_mock.dag_versions = [MagicMock(bundle_name="bundle_name")]
mock_get_team_name.return_value = "from_db"

assert DagRunInfo.team_name(dagrun_mock) == team_name

mock_get_team_name.assert_not_called()
mock_object_session.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.models.dagbundle.DagBundleModel.get_team_name")
@patch("airflow.models.dag.DagModel.get_team_name", side_effect=RuntimeError("db gone"))
@patch("sqlalchemy.orm.object_session")
@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True)
def test_dag_run_team_name_lookup_failure_does_not_raise(mock_getboolean, mock_get_team_name):
def test_dag_run_team_name_lookup_failure_does_not_raise(
mock_getboolean, mock_object_session, mock_get_team_name
):
"""A failed lookup must degrade to None -- `_cast_fields` would otherwise lose the whole event."""
dagrun_mock = MagicMock(DagRun)
dagrun_mock.dag_versions = [MagicMock(bundle_name="bundle_name")]
mock_get_team_name.side_effect = RuntimeError("Session must be set before!")
mock_object_session.return_value = MagicMock()
dagrun_mock = MagicMock(spec=DagRun)
dagrun_mock.dag_id = "test_dag"

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.3+")
@patch("airflow.models.dagbundle.DagBundleModel.get_team_name")
@patch("sqlalchemy.orm.object_session", return_value=None)
@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
def test_dag_run_team_name_no_session(mock_getboolean, mock_object_session):
"""When the dagrun has no attached session the lookup is skipped and None is returned."""
dagrun_mock = MagicMock(spec=DagRun)
dagrun_mock.dag_id = "test_dag"

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+")
Expand All @@ -430,6 +456,44 @@ def test_dag_run_team_name_multi_team_disabled(mock_getboolean, mock_get_team_na
mock_get_team_name.assert_not_called()


@pytest.mark.db_test
@patch("airflow.providers.openlineage.utils.utils.AIRFLOW_V_3_3_PLUS", False)
def test_dag_run_team_name_below_airflow_3_3():
"""Airflow < 3.3 has no multi-team support — team_name must return None unconditionally."""
dagrun_mock = MagicMock(spec=DagRun)

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.3+")
@patch("sqlalchemy.orm.object_session")
@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True)
def test_dag_run_team_name_from_scheduler_stamp(mock_getboolean, mock_object_session):
"""Scheduler stamps _team_name on ORM DagRun objects; the attribute is read back as-is.

The cascade must stop at `_team_name` — the DB lookup (object_session) must never be reached.
"""
dagrun_mock = MagicMock(spec=DagRun)
dagrun_mock._team_name = "team_a"

assert DagRunInfo.team_name(dagrun_mock) == "team_a"

mock_object_session.assert_not_called()


@pytest.mark.db_test
@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="multi-team requires Airflow 3.3+")
@patch("sqlalchemy.orm.object_session", return_value=None)
@patch("airflow.providers.openlineage.utils.utils.airflow_conf.getboolean", return_value=True)
def test_dag_run_team_name_from_scheduler_stamp_non_str(mock_getboolean, mock_object_session):
"""Non-str _team_name (corrupted stamp) must not propagate — fall through returns None."""
dagrun_mock = MagicMock(spec=DagRun)
dagrun_mock._team_name = 42

assert DagRunInfo.team_name(dagrun_mock) is None


def test_get_fully_qualified_class_name_serialized_operator():
op_module_path = BASH_OPERATOR_PATH
op_name = "BashOperator"
Expand Down Expand Up @@ -3025,6 +3089,7 @@ def test_dagrun_info_af3(mocked_dag_versions):
dv2.version_number = "version_number"
dv2.bundle_name = "bundle_name"
dv2.bundle_version = "bundle_version"
dv2.version_data = {"some": "data"}

optional_args = {}
if AIRFLOW_V_3_2_PLUS:
Expand Down Expand Up @@ -3059,12 +3124,16 @@ def test_dagrun_info_af3(mocked_dag_versions):
optional_result["partition_key"] = "some_partition_key"
optional_result["partition_date"] = "2024-06-01T00:00:00+00:00"

if AIRFLOW_V_3_3_PLUS:
optional_result["dag_version_data"] = {"some": "data"}

result = DagRunInfo(dagrun)
assert dict(result) == {
"conf": {"a": 1},
"clear_number": 0,
"dag_id": "dag_id",
"dag_team_name": None,
"dag_version_data": None,
"data_interval_end": "2024-06-01T00:00:00+00:00",
"data_interval_start": "2024-06-01T00:00:00+00:00",
"duration": 74.000546,
Expand Down Expand Up @@ -3127,6 +3196,7 @@ def test_dagrun_info_af2():
"dag_bundle_version": None,
"dag_version_id": None,
"dag_version_number": None,
"dag_version_data": None,
"note": None,
}

Expand Down Expand Up @@ -3169,6 +3239,7 @@ def test_taskinstance_info_af3():
assert dict(TaskInstanceInfo(runtime_ti)) == {
"log_url": runtime_ti.log_url,
"map_index": 2,
"note": None,
"rendered_map_index": None,
"try_number": 1,
"dag_bundle_version": "bundle_version",
Expand Down Expand Up @@ -3207,6 +3278,7 @@ def test_taskinstance_info_af2():
"log_url": "some_log_url",
"dag_bundle_name": None,
"dag_bundle_version": None,
"note": None,
}

# Also tested manually that it works well on AF2, hard to test hybrid property so just mocking it here
Expand Down