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 @@ -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

Expand Down Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 []
Expand All @@ -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,
Expand All @@ -213,6 +216,7 @@ def reset_metrics(worker_name: str) -> None:
sysinfo={
"status": logging.NOTSET,
},
team_name=team_name,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
WorkerApiDocs,
WorkerQueuesBody,
)
from airflow.utils.helpers import prune_dict
from airflow.utils.state import TaskInstanceState

if TYPE_CHECKING:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
WorkerSetStateReturn,
WorkerStateBody,
)
from airflow.utils.helpers import prune_dict

worker_router = AirflowRouter(
tags=["Worker"],
Expand Down Expand Up @@ -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(
Expand All @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down
63 changes: 61 additions & 2 deletions providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -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(
Expand Down Expand Up @@ -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")
Expand All @@ -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(
Expand Down Expand Up @@ -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")
Expand Down
Loading