diff --git a/providers/edge3/docs/deployment.rst b/providers/edge3/docs/deployment.rst index ebbdef2e5fb44..a48fea1bd53d7 100644 --- a/providers/edge3/docs/deployment.rst +++ b/providers/edge3/docs/deployment.rst @@ -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 diff --git a/providers/edge3/src/airflow/providers/edge3/cli/definition.py b/providers/edge3/src/airflow/providers/edge3/cli/definition.py index 97a609ad1a975..b1b89bbb46949 100644 --- a/providers/edge3/src/airflow/providers/edge3/cli/definition.py +++ b/providers/edge3/src/airflow/providers/edge3/cli/definition.py @@ -189,6 +189,7 @@ ARG_OUTPUT, ARG_STATE, ARG_WORKER_NAME_PATTERN, + ARG_QUEUES, ), ), ActionCommand( diff --git a/providers/edge3/src/airflow/providers/edge3/cli/edge_command.py b/providers/edge3/src/airflow/providers/edge3/cli/edge_command.py index 7a9186f2f574e..2088bd3bb7589 100644 --- a/providers/edge3/src/airflow/providers/edge3/cli/edge_command.py +++ b/providers/edge3/src/airflow/providers/edge3/cli/edge_command.py @@ -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 = [ 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 e20ca7db058f0..442b154593262 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/edge_worker.py +++ b/providers/edge3/src/airflow/providers/edge3/models/edge_worker.py @@ -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]: @@ -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 diff --git a/providers/edge3/tests/unit/edge3/cli/test_definition.py b/providers/edge3/tests/unit/edge3/cli/test_definition.py index cb225e89ffd92..cf6bb9e18ea7f 100644 --- a/providers/edge3/tests/unit/edge3/cli/test_definition.py +++ b/providers/edge3/tests/unit/edge3/cli/test_definition.py @@ -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 = [ diff --git a/providers/edge3/tests/unit/edge3/cli/test_worker.py b/providers/edge3/tests/unit/edge3/cli/test_worker.py index f959237e5adf4..2a7c2675bb78e 100644 --- a/providers/edge3/tests/unit/edge3/cli/test_worker.py +++ b/providers/edge3/tests/unit/edge3/cli/test_worker.py @@ -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: diff --git a/providers/edge3/tests/unit/edge3/models/test_edge_worker.py b/providers/edge3/tests/unit/edge3/models/test_edge_worker.py index 85ddefdb02e91..ab7c0550e7985 100644 --- a/providers/edge3/tests/unit/edge3/models/test_edge_worker.py +++ b/providers/edge3/tests/unit/edge3/models/test_edge_worker.py @@ -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): @@ -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"}