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
10 changes: 7 additions & 3 deletions airflow-core/src/airflow/assets/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@
)
from airflow.models.log import Log
from airflow.timetables.base import compute_rollup_fingerprint
from airflow.utils.helpers import is_container
from airflow.utils.helpers import is_container, prune_dict
from airflow.utils.log.logging_mixin import LoggingMixin
from airflow.utils.sqlalchemy import get_dialect_name, with_row_locks

Expand Down Expand Up @@ -397,15 +397,19 @@ def register_asset_change(
)
)

stats.incr("asset.updates")
team_name = None
if task_instance and conf.getboolean("core", "multi_team"):
from airflow.models.dag import DagModel

team_name = DagModel.get_team_name(task_instance.dag_id, session=session)
Comment thread
o-nikolas marked this conversation as resolved.
stats.incr("asset.updates", tags=prune_dict({"team_name": team_name}))

dags_to_queue = (
dags_to_queue_from_asset | dags_to_queue_from_asset_alias | dags_to_queue_from_asset_ref
)

if conf.getboolean("core", "multi_team"):
if task_instance:
team_name = DagModel.get_team_name(task_instance.dag_id, session=session)
resolved_source_teams = {team_name} if team_name else set()
# Resolve consumer-team filtering from the outlet reference
outlet_ref = session.scalar(
Expand Down
8 changes: 7 additions & 1 deletion airflow-core/src/airflow/jobs/scheduler_job_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,7 @@
from airflow.timetables.simple import AssetTriggeredTimetable
from airflow.triggers.base import TriggerEvent
from airflow.utils.event_scheduler import EventScheduler
from airflow.utils.helpers import prune_dict
from airflow.utils.log.logging_mixin import LoggingMixin
from airflow.utils.retries import MAX_DB_RETRIES, retry_db_transaction, run_with_db_retries
from airflow.utils.session import NEW_SESSION, create_session, provide_session
Expand Down Expand Up @@ -2585,7 +2586,12 @@ def _create_dag_runs_asset_triggered(
creating_job_id=self.job.id,
session=session,
)
stats.incr("asset.triggered_dagruns")
team_name = (
self._get_team_names_for_dag_ids([dag.dag_id], session).get(dag.dag_id)
if self._multi_team
else None
)
stats.incr("asset.triggered_dagruns", tags=prune_dict({"team_name": team_name}))
dag_run.consumed_asset_events.extend(asset_events)
self.log.info(
"Created asset-triggered DagRun for '%s': run_id=%s, consumed %d asset events",
Expand Down
57 changes: 56 additions & 1 deletion airflow-core/tests/unit/assets/test_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from sqlalchemy.orm import Session

from airflow import settings
from airflow._shared.observability.metrics.base_stats_logger import StatsLogger
from airflow.assets.manager import AssetManager
from airflow.models.asset import (
AssetAliasModel,
Expand All @@ -40,15 +41,21 @@
DagScheduleAssetReference,
)
from airflow.models.dag import DAG, DagModel
from airflow.models.dagbundle import DagBundleModel
from airflow.models.log import Log
from airflow.models.team import Team
from airflow.partition_mappers.temporal import FanOutMapper, StartOfWeekMapper
from airflow.partition_mappers.window import WeekWindow
from airflow.providers.standard.operators.empty import EmptyOperator
from airflow.sdk.definitions.asset import Asset
from airflow.sdk.definitions.timetables.assets import PartitionedAssetTimetable

from tests_common.test_utils.config import conf_vars
from tests_common.test_utils.db import clear_db_apdr, clear_db_logs, clear_db_pakl
from tests_common.test_utils.db import (
clear_db_apdr,
clear_db_logs,
clear_db_pakl,
)
from unit.listeners import asset_listener

pytestmark = pytest.mark.db_test
Expand Down Expand Up @@ -670,6 +677,54 @@ def _make_asset_model(
return model


class TestAssetMetricsTeamName:
@pytest.mark.parametrize(
("multi_team", "expect_team_tag"),
[
pytest.param("true", True, id="with_team"),
pytest.param("false", False, id="without_team"),
],
)
@mock.patch("airflow._shared.observability.metrics.stats._get_backend")
def test_asset_updates_respects_team_name(
self, mock_get_backend, multi_team, expect_team_tag, session, dag_maker
):
mock_stats = mock.MagicMock(spec=StatsLogger)
mock_get_backend.return_value = mock_stats

suffix = "with_team" if expect_team_tag else "without_team"

team_name = f"team_asset_upd_{suffix}"
team = Team(name=team_name)
session.add(team)
session.flush()

bundle_name = f"bundle_asset_upd_{suffix}"
bundle = DagBundleModel(name=bundle_name)
bundle.teams.append(team)
session.add(bundle)
session.flush()

asset_name = f"metric_asset_{suffix}"
asset = Asset(uri=f"test://{asset_name}", name=asset_name, group="asset")
with dag_maker(dag_id=f"asset_dag_{suffix}", bundle_name=bundle_name, session=session):
EmptyOperator(task_id="task1", outlets=[asset])

ti = mock.MagicMock()
ti.dag_id = f"asset_dag_{suffix}"
ti.task_id = "task1"
ti.run_id = "run1"
ti.map_index = -1

with conf_vars({("core", "multi_team"): multi_team}):
AssetManager().register_asset_change(task_instance=ti, asset=asset, session=session)

if expect_team_tag:
mock_stats.incr.assert_any_call("asset.updates", tags={"team_name": team_name})
else:
mock_stats.incr.assert_any_call("asset.updates")


class TestFilterDagsByTeam:
@conf_vars({("core", "multi_team"): "false"})
def test_multi_team_disabled_returns_all_dags(self):
Expand Down
61 changes: 61 additions & 0 deletions airflow-core/tests/unit/jobs/test_scheduler_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -5503,6 +5503,67 @@ def dict_from_obj(obj):
assert created_run.data_interval_end is None
assert created_run.creating_job_id == scheduler_job.id

@pytest.mark.parametrize(
("multi_team", "expect_team_tag"),
[
pytest.param("true", True, id="with_team"),
pytest.param("false", False, id="without_team"),
],
)
@mock.patch("airflow._shared.observability.metrics.stats._get_backend")
def test_asset_triggered_dagruns_respects_team_name(
self, mock_get_backend, multi_team, expect_team_tag, session, dag_maker
):
mock_stats = mock.MagicMock(spec=StatsLogger)
mock_get_backend.return_value = mock_stats

suffix = "with_team" if expect_team_tag else "without_team"

team_name = f"team_asset_trig_{suffix}"
team = Team(name=team_name)
session.add(team)
session.flush()

bundle_name = f"bundle_asset_trig_{suffix}"
bundle = DagBundleModel(name=bundle_name)
bundle.teams.append(team)
session.add(bundle)
session.commit()

asset_name = f"test_team_asset_{suffix}"
asset = Asset(uri=f"test://{asset_name}", name=asset_name, group="test_group")
with dag_maker(dag_id=f"producer_{suffix}", bundle_name=bundle_name, session=session):
BashOperator(task_id="task", bash_command="echo 1", outlets=[asset])
dr = dag_maker.create_dagrun()

asset_id = session.scalar(select(AssetModel.id).where(AssetModel.uri == asset.uri))
event = AssetEvent(
asset_id=asset_id,
source_task_id="task",
source_dag_id=dr.dag_id,
source_run_id=dr.run_id,
source_map_index=-1,
)
session.add(event)

with dag_maker(
dag_id=f"consumer_{suffix}", schedule=[asset], bundle_name=bundle_name, session=session
):
pass

session.add(AssetDagRunQueue(asset_id=asset_id, target_dag_id=f"consumer_{suffix}"))
session.flush()

with conf_vars({("core", "multi_team"): multi_team}):
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(job=scheduler_job, executors=[self.null_exec])
self.job_runner._create_dagruns_for_dags(session, session)

if expect_team_tag:
mock_stats.incr.assert_any_call("asset.triggered_dagruns", tags={"team_name": team_name})
else:
mock_stats.incr.assert_any_call("asset.triggered_dagruns")

@pytest.mark.need_serialized_dag
@pytest.mark.parametrize(
("disable", "enable"),
Expand Down
Loading