From 74f42fcc22dbd8330fcd3ffc38eda4249c8b8fb6 Mon Sep 17 00:00:00 2001 From: Shivam Rastogi <6463385+shivaam@users.noreply.github.com> Date: Thu, 25 Jun 2026 16:28:02 -0700 Subject: [PATCH] Add team name tags to Edge worker metrics --- .../edge3/executors/edge_executor.py | 5 +- .../providers/edge3/models/edge_worker.py | 22 ++++--- .../providers/edge3/worker_api/routes/jobs.py | 8 ++- .../edge3/worker_api/routes/worker.py | 9 ++- .../edge3/executors/test_edge_executor.py | 42 +++++++++--- .../unit/edge3/worker_api/routes/test_jobs.py | 63 +++++++++++++++++- .../edge3/worker_api/routes/test_worker.py | 64 ++++++++++++++++++- 7 files changed, 187 insertions(+), 26 deletions(-) diff --git a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py index 5c582781a45f0..e6f7af33eb06c 100644 --- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py +++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py @@ -219,7 +219,7 @@ def _check_worker_liveness(self, session: Session) -> bool: sysinfo.pop("status_text", None) # Remove old status text if exists worker.sysinfo = sysinfo self.log.warning("Worker %s is lifeless. Setting state to %s", worker.worker_name, worker.state) - reset_metrics(worker.worker_name) + reset_metrics(worker.worker_name, team_name=worker.team_name) return changed @@ -254,8 +254,9 @@ def _update_orphaned_jobs(self, session: Session) -> bool: "task_id": job.task_id, "queue": job.queue, "state": str(TaskInstanceState.FAILED), + "team_name": job.team_name, } - Stats.incr("edge_worker.ti.finish", tags=tags) + Stats.incr("edge_worker.ti.finish", tags=prune_dict(tags)) return bool(lifeless_jobs) diff --git a/providers/edge3/src/airflow/providers/edge3/models/edge_worker.py b/providers/edge3/src/airflow/providers/edge3/models/edge_worker.py index 78c60307ea1cd..fb800b8b2457b 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/edge_worker.py +++ b/providers/edge3/src/airflow/providers/edge3/models/edge_worker.py @@ -28,6 +28,7 @@ from airflow.providers.common.compat.sdk import AirflowException, Stats, timezone from airflow.providers.common.compat.sqlalchemy.orm import mapped_column from airflow.providers.edge3.models.edge_base import Base +from airflow.utils.helpers import prune_dict from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.providers_configuration_loader import providers_configuration_loaded from airflow.utils.session import NEW_SESSION, provide_session @@ -156,6 +157,7 @@ def set_metrics( free_concurrency: int, queues: list[str] | None, sysinfo: dict[str, str | int | float | datetime], + team_name: str | None = None, ) -> None: """Set metric of edge worker.""" queues = queues if queues else [] @@ -178,30 +180,31 @@ def set_metrics( "concurrency", "free_concurrency", } + metric_tags = prune_dict({"worker_name": worker_name, "team_name": team_name}) Stats.gauge( "edge_worker.status", sysinfo.get("status", logging.NOTSET), # type: ignore - tags={"worker_name": worker_name}, + tags=metric_tags, ) - Stats.gauge("edge_worker.connected", int(connected), tags={"worker_name": worker_name}) - Stats.gauge("edge_worker.maintenance", int(maintenance), tags={"worker_name": worker_name}) - Stats.gauge("edge_worker.jobs_active", jobs_active, tags={"worker_name": worker_name}) - Stats.gauge("edge_worker.concurrency", concurrency, tags={"worker_name": worker_name}) - Stats.gauge("edge_worker.free_concurrency", free_concurrency, tags={"worker_name": worker_name}) + Stats.gauge("edge_worker.connected", int(connected), tags=metric_tags) + Stats.gauge("edge_worker.maintenance", int(maintenance), tags=metric_tags) + Stats.gauge("edge_worker.jobs_active", jobs_active, tags=metric_tags) + Stats.gauge("edge_worker.concurrency", concurrency, tags=metric_tags) + Stats.gauge("edge_worker.free_concurrency", free_concurrency, tags=metric_tags) Stats.gauge( "edge_worker.num_queues", len(queues), - tags={"worker_name": worker_name, "queues": ",".join(queues)}, + tags={**metric_tags, "queues": ",".join(queues)}, ) for key in additional_keys: value = sysinfo.get(key) if isinstance(value, (int, float)): - Stats.gauge(f"edge_worker.{key}", value, tags={"worker_name": worker_name}) + Stats.gauge(f"edge_worker.{key}", value, tags=metric_tags) -def reset_metrics(worker_name: str) -> None: +def reset_metrics(worker_name: str, team_name: str | None = None) -> None: """Reset metrics of worker.""" set_metrics( worker_name=worker_name, @@ -213,6 +216,7 @@ def reset_metrics(worker_name: str) -> None: sysinfo={ "status": logging.NOTSET, }, + team_name=team_name, ) diff --git a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py index 7ed84b180d771..cf992b7dc8deb 100644 --- a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py +++ b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py @@ -36,6 +36,7 @@ WorkerApiDocs, WorkerQueuesBody, ) +from airflow.utils.helpers import prune_dict from airflow.utils.state import TaskInstanceState if TYPE_CHECKING: @@ -104,7 +105,9 @@ def fetch( job.last_update = timezone.utcnow() session.commit() # Edge worker does not backport emitted Airflow metrics, so export some metrics - tags = {"dag_id": job.dag_id, "task_id": job.task_id, "queue": job.queue} + tags = prune_dict( + {"dag_id": job.dag_id, "task_id": job.task_id, "queue": job.queue, "team_name": job.team_name} + ) Stats.incr("edge_worker.ti.start", tags=tags) return EdgeJobFetched( dag_id=job.dag_id, @@ -157,8 +160,9 @@ def state( "task_id": job.task_id, "queue": job.queue, "state": str(state), + "team_name": job.team_name, } - Stats.incr("edge_worker.ti.finish", tags=tags) + Stats.incr("edge_worker.ti.finish", tags=prune_dict(tags)) query2 = ( update(EdgeJobModel) diff --git a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/worker.py b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/worker.py index 6ca8e794e9953..2e42ae8902b60 100644 --- a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/worker.py +++ b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/worker.py @@ -38,6 +38,7 @@ WorkerSetStateReturn, WorkerStateBody, ) +from airflow.utils.helpers import prune_dict worker_router = AirflowRouter( tags=["Worker"], @@ -244,7 +245,12 @@ def set_state( worker.sysinfo = body.sysinfo worker.last_update = timezone.utcnow() session.commit() - Stats.incr("edge_worker.heartbeat_count", 1, 1, tags={"worker_name": worker_name}) + Stats.incr( + "edge_worker.heartbeat_count", + 1, + 1, + tags=prune_dict({"worker_name": worker_name, "team_name": worker.team_name}), + ) concurrency: int = body.sysinfo.get("concurrency", -1) # type: ignore free_concurrency: int = body.sysinfo.get("free_concurrency", -1) # type: ignore set_metrics( @@ -255,6 +261,7 @@ def set_state( free_concurrency=free_concurrency, queues=worker.queues, sysinfo=body.sysinfo, + team_name=worker.team_name, ) versions_match = _assert_version(body.sysinfo) # Exception only after worker state is in the DB return WorkerSetStateReturn( diff --git a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py index 8975eb7ad9443..2f135957c4743 100644 --- a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py +++ b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py @@ -67,9 +67,23 @@ def get_test_executor(self, pool_slots=1): return (executor, key) + @pytest.mark.parametrize( + ("executor_kwargs", "job_team_name", "expected_tags"), + [ + ({}, None, {}), + pytest.param( + {"team_name": "team_a"}, + "team_a", + {"team_name": "team_a"}, + marks=pytest.mark.skipif( + not AIRFLOW_V_3_2_PLUS, reason="team_name is only available in Airflow 3.2+" + ), + ), + ], + ) @patch(f"{Stats.__module__}.Stats.incr") - def test_sync_orphaned_tasks(self, mock_stats_incr): - executor = EdgeExecutor() + def test_sync_orphaned_tasks(self, mock_stats_incr, executor_kwargs, job_team_name, expected_tags): + executor = EdgeExecutor(**executor_kwargs) delta_to_purge = timedelta(minutes=conf.getint("edge", "job_fail_purge") + 1) delta_to_orphaned_config_name = "task_instance_heartbeat_timeout" @@ -97,20 +111,23 @@ def test_sync_orphaned_tasks(self, mock_stats_incr): command="mock", concurrency_slots=1, last_update=last_update, + team_name=job_team_name, ) ) session.commit() + expected_tags = { + "dag_id": "test_dag", + "queue": "default", + "state": "failed", + "task_id": "started_running_orphaned", + **expected_tags, + } executor.sync() mock_stats_incr.assert_called_with( "edge_worker.ti.finish", - tags={ - "dag_id": "test_dag", - "queue": "default", - "state": "failed", - "task_id": "started_running_orphaned", - }, + tags=expected_tags, ) assert mock_stats_incr.call_count == 1 @@ -549,10 +566,17 @@ def test_check_worker_liveness_filters_by_team_name(self): with time_machine.travel(datetime(2023, 1, 1, 1, 0, 0, tzinfo=timezone.utc), tick=False): with conf_vars({("edge", "heartbeat_interval"): "10"}): - with create_session() as session: + with ( + create_session() as session, + patch( + "airflow.providers.edge3.executors.edge_executor.reset_metrics" + ) as mock_reset_metrics, + ): executor_a._check_worker_liveness(session) session.commit() + mock_reset_metrics.assert_called_once_with("worker_team_a", team_name="team_a") + with create_session() as session: workers = {w.worker_name: w for w in session.scalars(select(EdgeWorkerModel)).all()} assert workers["worker_team_a"].state == EdgeWorkerState.UNKNOWN diff --git a/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py b/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py index 26a06ba470f4e..0b09e47b633f8 100644 --- a/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py +++ b/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py @@ -98,6 +98,7 @@ def test_state(self, mock_stats_incr, session: Session): queue=QUEUE, concurrency_slots=1, command="execute", + team_name="team_a", ) session.add(job) session.commit() @@ -130,6 +131,7 @@ def test_state(self, mock_stats_incr, session: Session): "queue": QUEUE, "state": TaskInstanceState.SUCCESS, "task_id": TASK_ID, + "team_name": "team_a", }, ) assert mock_stats_incr.call_count == 1 @@ -138,7 +140,45 @@ def test_state(self, mock_stats_incr, session: Session): assert db_job is not None assert db_job.state == TaskInstanceState.SUCCESS - def test_fetch_filters_by_worker_team_name(self, session: Session): + @patch(f"{Stats.__module__}.Stats.incr") + def test_state_finish_metric_omits_team_name_for_global_job(self, mock_stats_incr, session: Session): + with create_session() as session: + job = EdgeJobModel( + dag_id=DAG_ID, + task_id=TASK_ID, + run_id=RUN_ID, + try_number=1, + map_index=-1, + state=TaskInstanceState.RUNNING, + queue=QUEUE, + concurrency_slots=1, + command="execute", + ) + session.add(job) + session.commit() + + state( + dag_id=DAG_ID, + task_id=TASK_ID, + run_id=RUN_ID, + try_number=1, + map_index=-1, + state=TaskInstanceState.SUCCESS, + session=session, + ) + + mock_stats_incr.assert_called_once_with( + "edge_worker.ti.finish", + tags={ + "dag_id": DAG_ID, + "queue": QUEUE, + "state": TaskInstanceState.SUCCESS, + "task_id": TASK_ID, + }, + ) + + @patch(f"{Stats.__module__}.Stats.incr") + def test_fetch_filters_by_worker_team_name(self, mock_stats_incr, session: Session): with create_session() as session: session.add( EdgeWorkerModel( @@ -177,6 +217,15 @@ def test_fetch_filters_by_worker_team_name(self, session: Session): assert result is not None assert result.dag_id == "dag_a" assert result.task_id == "task_a" + mock_stats_incr.assert_called_once_with( + "edge_worker.ti.start", + tags={ + "dag_id": "dag_a", + "queue": QUEUE, + "task_id": "task_a", + "team_name": "team_a", + }, + ) def test_fetch_unknown_worker_raises_404(self, session: Session): body = WorkerQueuesBody(free_concurrency=1, queues=[QUEUE], team_name="team_a") @@ -187,7 +236,8 @@ def test_fetch_unknown_worker_raises_404(self, session: Session): assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND assert exc_info.value.detail == "Worker not found" - def test_fetch_without_team_name_returns_any_team(self, session: Session): + @patch(f"{Stats.__module__}.Stats.incr") + def test_fetch_without_team_name_returns_any_team(self, mock_stats_incr, session: Session): """When a worker has no team_name, no team filter is applied so any queued job can be returned.""" with create_session() as session: session.add( @@ -232,6 +282,15 @@ def test_fetch_without_team_name_returns_any_team(self, session: Session): assert result3 is None fetched_dag_ids = {result1.dag_id, result2.dag_id} assert fetched_dag_ids == {"dag_a", "dag_b"} + mock_stats_incr.assert_any_call( + "edge_worker.ti.start", + tags={"dag_id": "dag_a", "queue": QUEUE, "task_id": "task_a", "team_name": "team_a"}, + ) + mock_stats_incr.assert_any_call( + "edge_worker.ti.start", + tags={"dag_id": "dag_b", "queue": QUEUE, "task_id": "task_b"}, + ) + assert mock_stats_incr.call_count == 2 @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="The tests should be skipped for Airflow < 3.3") diff --git a/providers/edge3/tests/unit/edge3/worker_api/routes/test_worker.py b/providers/edge3/tests/unit/edge3/worker_api/routes/test_worker.py index 099ba7e882739..f42d93950f3b7 100644 --- a/providers/edge3/tests/unit/edge3/worker_api/routes/test_worker.py +++ b/providers/edge3/tests/unit/edge3/worker_api/routes/test_worker.py @@ -20,13 +20,14 @@ from datetime import datetime from pathlib import Path from typing import TYPE_CHECKING +from unittest.mock import patch import pytest from fastapi import HTTPException from sqlalchemy import delete, select from airflow import __version__ as airflow_version -from airflow.providers.common.compat.sdk import timezone +from airflow.providers.common.compat.sdk import Stats, timezone from airflow.providers.edge3 import __version__ as edge_provider_version from airflow.providers.edge3.cli.worker import EdgeWorker from airflow.providers.edge3.models.edge_worker import ( @@ -363,6 +364,67 @@ def test_set_state(self, session: Session, cli_worker: EdgeWorker): assert worker[0].queues == queues assert return_queues == ["default", "default2"] + @pytest.mark.parametrize( + ("worker_team_name", "expected_worker_tags"), + [ + pytest.param("team_a", {"worker_name": "test2_worker", "team_name": "team_a"}, id="team"), + pytest.param(None, {"worker_name": "test2_worker"}, id="global"), + ], + ) + @patch(f"{Stats.__module__}.Stats.gauge") + @patch(f"{Stats.__module__}.Stats.incr") + def test_set_state_metrics_team_name_tags( + self, + mock_stats_incr, + mock_stats_gauge, + session: Session, + cli_worker: EdgeWorker, + worker_team_name: str | None, + expected_worker_tags: dict[str, str], + ): + queues = ["default", "default2"] + rwm = EdgeWorkerModel( + worker_name="test2_worker", + state=EdgeWorkerState.IDLE, + queues=queues, + first_online=timezone.utcnow(), + team_name=worker_team_name, + ) + session.add(rwm) + session.commit() + + body = WorkerStateBody( + state=EdgeWorkerState.RUNNING, + jobs_active=1, + queues=["default2"], + sysinfo={**self.MOCK_SYSINFO, "disk_usage": 42.5, "status_text": "ok"}, + ) + set_state("test2_worker", body, session) + + mock_stats_incr.assert_called_once_with( + "edge_worker.heartbeat_count", + 1, + 1, + tags=expected_worker_tags, + ) + mock_stats_gauge.assert_any_call( + "edge_worker.status", + self.MOCK_SYSINFO["status"], + tags=expected_worker_tags, + ) + mock_stats_gauge.assert_any_call("edge_worker.connected", 1, tags=expected_worker_tags) + mock_stats_gauge.assert_any_call("edge_worker.maintenance", 0, tags=expected_worker_tags) + mock_stats_gauge.assert_any_call("edge_worker.jobs_active", 1, tags=expected_worker_tags) + mock_stats_gauge.assert_any_call("edge_worker.concurrency", 8, tags=expected_worker_tags) + mock_stats_gauge.assert_any_call("edge_worker.free_concurrency", 8, tags=expected_worker_tags) + mock_stats_gauge.assert_any_call( + "edge_worker.num_queues", + len(queues), + tags={**expected_worker_tags, "queues": ",".join(queues)}, + ) + mock_stats_gauge.assert_any_call("edge_worker.disk_usage", 42.5, tags=expected_worker_tags) + assert mock_stats_gauge.call_count == 8 + def test_set_state_returns_concurrency(self, session: Session, cli_worker: EdgeWorker): """set_state includes the DB-stored concurrency override in its response.""" rwm = EdgeWorkerModel(