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
3 changes: 2 additions & 1 deletion providers/edge3/docs/deployment.rst
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,8 @@ instance. The commands are:

- ``airflow edge list-workers``: List all workers in the cluster. Accepts an optional
``--worker-name-pattern`` glob (e.g. ``'prod-*'``) to filter workers by name,
and ``-s``/``--state`` to filter by worker state.
``-s``/``--state`` to filter by worker state, and ``-q``/``--queues`` (comma
delimited) to list only workers serving any of the given queues.
- ``airflow edge remote-edge-worker-request-maintenance``: Request a remote edge worker to enter maintenance mode
- ``airflow edge remote-edge-worker-update-maintenance-comment``: Updates the maintenance comment for a remote edge worker
- ``airflow edge remote-edge-worker-exit-maintenance``: Request a remote edge worker to exit maintenance mode
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -189,6 +189,7 @@
ARG_OUTPUT,
ARG_STATE,
ARG_WORKER_NAME_PATTERN,
ARG_QUEUES,
),
),
ActionCommand(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -261,7 +261,9 @@ def list_edge_workers(args) -> None:
from airflow.providers.edge3.models.edge_worker import get_registered_edge_hosts

all_hosts_iter = get_registered_edge_hosts(
states=args.state, worker_name_pattern=args.worker_name_pattern
states=args.state,
worker_name_pattern=args.worker_name_pattern,
queues=args.queues.split(",") if args.queues else None,
)
# Format and print worker info on the screen
fields = [
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -256,6 +256,7 @@ def _fetch_edge_hosts_from_db(
hostname: str | None = None,
states: list | None = None,
worker_name_pattern: str | None = None,
queues: list[str] | None = None,
*,
session: Session = NEW_SESSION,
) -> Sequence[EdgeWorkerModel]:
Expand All @@ -269,15 +270,28 @@ def _fetch_edge_hosts_from_db(
EdgeWorkerModel.worker_name.like(_glob_to_like_pattern(worker_name_pattern), escape="\\")
)
query = query.order_by(EdgeWorkerModel.worker_name)
return session.scalars(query).all()
workers = session.scalars(query).all()
if queues:
# Queues are stored as a repr-encoded list in a single column, so exact
# membership is filtered in Python to avoid substring false positives. A
# worker matches if it serves any of the requested queues.
wanted = set(queues)
workers = [worker for worker in workers if worker.queues and wanted.intersection(worker.queues)]
return workers


@providers_configuration_loaded
@provide_session
def get_registered_edge_hosts(
*, states: list | None = None, worker_name_pattern: str | None = None, session: Session = NEW_SESSION
*,
states: list | None = None,
worker_name_pattern: str | None = None,
queues: list[str] | None = None,
session: Session = NEW_SESSION,
):
return _fetch_edge_hosts_from_db(states=states, worker_name_pattern=worker_name_pattern, session=session)
return _fetch_edge_hosts_from_db(
states=states, worker_name_pattern=worker_name_pattern, queues=queues, session=session
)


@provide_session
Expand Down
6 changes: 6 additions & 0 deletions providers/edge3/tests/unit/edge3/cli/test_definition.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,12 @@ def test_list_workers_command_args(self):
assert args.state == ["running", "maintenance"]
assert args.worker_name_pattern == "prod-*"

def test_list_workers_command_queues_arg(self):
"""Test list-workers command with the queues filter."""
params = ["edge", "list-workers", "--queues", "gpu,default"]
args = self.arg_parser.parse_args(params)
assert args.queues == "gpu,default"

def test_remote_edge_worker_request_maintenance_args(self):
"""Test remote-edge-worker-request-maintenance command with required arguments."""
params = [
Expand Down
20 changes: 19 additions & 1 deletion providers/edge3/tests/unit/edge3/cli/test_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1130,7 +1130,25 @@ def test_list_edge_workers_passes_name_pattern(self, mock_edgeworker: EdgeWorker
) as mock_get_hosts,
):
edge_command.list_edge_workers(args)
mock_get_hosts.assert_called_once_with(states=None, worker_name_pattern="prod-*")
mock_get_hosts.assert_called_once_with(states=None, worker_name_pattern="prod-*", queues=None)

@pytest.mark.db_test
def test_list_edge_workers_passes_queues(self, mock_edgeworker: EdgeWorkerModel):
args = self.parser.parse_args(["edge", "list-workers", "--output", "json", "--queues", "gpu,default"])
with contextlib.redirect_stdout(StringIO()):
with (
patch(
"airflow.providers.edge3.cli.edge_command._check_valid_db_connection",
),
patch(
"airflow.providers.edge3.models.edge_worker.get_registered_edge_hosts",
return_value=[mock_edgeworker],
) as mock_get_hosts,
):
edge_command.list_edge_workers(args)
mock_get_hosts.assert_called_once_with(
states=None, worker_name_pattern=None, queues=["gpu", "default"]
)


class TestSignalHandling:
Expand Down
25 changes: 23 additions & 2 deletions providers/edge3/tests/unit/edge3/models/test_edge_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,8 +97,13 @@ class TestGetRegisteredEdgeHosts:
@pytest.fixture(autouse=True)
def setup_test_cases(self, session: Session):
session.execute(delete(EdgeWorkerModel))
for name in ("prod-worker-1", "prod-worker-2", "dev-worker-1"):
session.add(EdgeWorkerModel(worker_name=name, queues=["default"], state=EdgeWorkerState.RUNNING))
queues_by_name = {
"prod-worker-1": ["default", "gpu"],
"prod-worker-2": ["default"],
"dev-worker-1": ["gpu"],
}
for name, queues in queues_by_name.items():
session.add(EdgeWorkerModel(worker_name=name, queues=queues, state=EdgeWorkerState.RUNNING))
session.commit()

def test_no_pattern_returns_all(self, session: Session):
Expand All @@ -120,3 +125,19 @@ def test_question_mark_glob_matches_single_char(self, session: Session):
def test_no_match_returns_empty(self, session: Session):
hosts = get_registered_edge_hosts(worker_name_pattern="nonexistent-*", session=session)
assert list(hosts) == []

def test_queues_filters_by_exact_membership(self, session: Session):
hosts = get_registered_edge_hosts(queues=["gpu"], session=session)
assert {h.worker_name for h in hosts} == {"prod-worker-1", "dev-worker-1"}

def test_queues_matches_any_of_multiple(self, session: Session):
hosts = get_registered_edge_hosts(queues=["gpu", "default"], session=session)
assert {h.worker_name for h in hosts} == {"prod-worker-1", "prod-worker-2", "dev-worker-1"}

def test_queues_no_match_returns_empty(self, session: Session):
hosts = get_registered_edge_hosts(queues=["nonexistent"], session=session)
assert list(hosts) == []

def test_queues_combined_with_name_pattern(self, session: Session):
hosts = get_registered_edge_hosts(worker_name_pattern="prod-*", queues=["gpu"], session=session)
assert {h.worker_name for h in hosts} == {"prod-worker-1"}