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 @@ -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
Expand Down Expand Up @@ -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 AirflowDagRunFacet, AirflowRunFacet
from airflow.sdk.execution_time.secrets_masker import SecretsMasker, _secrets_masker
from airflow.utils.state import DagRunState
else:
Expand Down Expand Up @@ -172,10 +174,33 @@ 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"))
Comment thread
SameerMesiah97 marked this conversation as resolved.

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(
"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(
Expand All @@ -184,7 +209,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(),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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:
Expand Down Expand Up @@ -268,9 +270,18 @@ def on_running():
if not doc:
doc, doc_type = get_dag_documentation(dag)

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,
Expand All @@ -284,6 +295,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(
Expand Down Expand Up @@ -318,10 +330,17 @@ def on_running():
},
)
event_size = len(Serde.to_json(redacted_event).encode("utf-8"))

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)
Expand Down Expand Up @@ -413,9 +432,18 @@ def on_success():
if not doc:
doc, doc_type = get_dag_documentation(dag)

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,
Expand All @@ -429,6 +457,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(
Expand Down Expand Up @@ -462,10 +491,17 @@ def on_success():
},
)
event_size = len(Serde.to_json(redacted_event).encode("utf-8"))

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)
Expand Down Expand Up @@ -572,9 +608,18 @@ def on_failure():
if not doc:
doc, doc_type = get_dag_documentation(dag)

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,
Expand Down Expand Up @@ -622,10 +667,17 @@ def on_failure():
},
)
event_size = len(Serde.to_json(redacted_event).encode("utf-8"))

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)
Expand Down Expand Up @@ -708,9 +760,18 @@ def on_skipped():
if not doc:
doc, doc_type = get_dag_documentation(dag)

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,
Expand Down Expand Up @@ -757,10 +818,17 @@ def on_skipped():
},
)
event_size = len(Serde.to_json(redacted_event).encode("utf-8"))

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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@
BaseOperator,
BaseSensorOperator,
MappedOperator,
conf as airflow_conf,
)
from airflow.providers.openlineage import (
__version__ as OPENLINEAGE_PROVIDER_VERSION,
Expand All @@ -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
Expand Down Expand Up @@ -980,6 +982,7 @@ 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),
}

Expand Down Expand Up @@ -1037,7 +1040,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:
Expand All @@ -1053,6 +1056,20 @@ 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

from airflow.models.dagbundle import DagBundleModel

bundle_name = cls.dag_version_info(dagrun, "bundle_name")
Comment thread
SameerMesiah97 marked this conversation as resolved.
if not isinstance(bundle_name, str):
return None

return DagBundleModel.get_team_name(bundle_name)


class TaskInstanceInfo(InfoJsonEncodable):
"""Defines encoding TaskInstance object to JSON."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Loading