From 625c1297a701c931f63e82bb56f711d6bb0268ac Mon Sep 17 00:00:00 2001 From: shubhamraj-git Date: Sat, 25 Jul 2026 13:53:18 +0000 Subject: [PATCH 1/2] add queue paarmeter to list-workers --- providers/edge3/docs/deployment.rst | 3 ++- .../airflow/providers/edge3/cli/definition.py | 6 ++++++ .../providers/edge3/cli/edge_command.py | 2 +- .../providers/edge3/models/edge_worker.py | 18 +++++++++++++--- .../tests/unit/edge3/cli/test_definition.py | 6 ++++++ .../edge3/tests/unit/edge3/cli/test_worker.py | 18 +++++++++++++++- .../unit/edge3/models/test_edge_worker.py | 21 +++++++++++++++++-- 7 files changed, 66 insertions(+), 8 deletions(-) diff --git a/providers/edge3/docs/deployment.rst b/providers/edge3/docs/deployment.rst index ebbdef2e5fb44..d5e37d6041b4b 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``/``--queue`` to list only + workers serving a given queue. - ``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..ad15af57ac049 100644 --- a/providers/edge3/src/airflow/providers/edge3/cli/definition.py +++ b/providers/edge3/src/airflow/providers/edge3/cli/definition.py @@ -57,6 +57,11 @@ metavar="PATTERN", help="Optional glob pattern to filter workers by name, e.g. 'prod-*'. Lists all workers if omitted.", ) +ARG_QUEUE = Arg( + ("-q", "--queue"), + metavar="QUEUE", + help="Optional queue name to filter workers by. Only workers serving this exact queue are listed.", +) ARG_REQUIRED_EDGE_HOSTNAME = Arg( ("-H", "--edge-hostname"), help="Set the hostname of worker if you have multiple workers on a single machine", @@ -189,6 +194,7 @@ ARG_OUTPUT, ARG_STATE, ARG_WORKER_NAME_PATTERN, + ARG_QUEUE, ), ), 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..5bdb9dc82cbbd 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,7 @@ 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, queue=args.queue ) # 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..f36d662ce8216 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, + queue: str | None = None, *, session: Session = NEW_SESSION, ) -> Sequence[EdgeWorkerModel]: @@ -269,15 +270,26 @@ 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 queue: + # Queues are stored as a repr-encoded list in a single column, so exact + # membership is filtered in Python to avoid substring false positives. + workers = [worker for worker in workers if worker.queues and queue in 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, + queue: 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, queue=queue, 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..2cdaa38731341 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_queue_arg(self): + """Test list-workers command with the queue filter.""" + params = ["edge", "list-workers", "--queue", "gpu"] + args = self.arg_parser.parse_args(params) + assert args.queue == "gpu" + 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..5b075e29fb1d8 100644 --- a/providers/edge3/tests/unit/edge3/cli/test_worker.py +++ b/providers/edge3/tests/unit/edge3/cli/test_worker.py @@ -1130,7 +1130,23 @@ 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-*", queue=None) + + @pytest.mark.db_test + def test_list_edge_workers_passes_queue(self, mock_edgeworker: EdgeWorkerModel): + args = self.parser.parse_args(["edge", "list-workers", "--output", "json", "--queue", "gpu"]) + 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, queue="gpu") 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..e246d592d53f8 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,15 @@ 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_queue_filters_by_exact_membership(self, session: Session): + hosts = get_registered_edge_hosts(queue="gpu", session=session) + assert {h.worker_name for h in hosts} == {"prod-worker-1", "dev-worker-1"} + + def test_queue_no_match_returns_empty(self, session: Session): + hosts = get_registered_edge_hosts(queue="nonexistent", session=session) + assert list(hosts) == [] + + def test_queue_combined_with_name_pattern(self, session: Session): + hosts = get_registered_edge_hosts(worker_name_pattern="prod-*", queue="gpu", session=session) + assert {h.worker_name for h in hosts} == {"prod-worker-1"} From 6c4d4fd7a29044ad251a99bfc86fd7d54c32a49b Mon Sep 17 00:00:00 2001 From: shubhamraj-git Date: Sun, 26 Jul 2026 06:36:04 +0000 Subject: [PATCH 2/2] refactor --- providers/edge3/docs/deployment.rst | 4 ++-- .../airflow/providers/edge3/cli/definition.py | 7 +------ .../airflow/providers/edge3/cli/edge_command.py | 4 +++- .../providers/edge3/models/edge_worker.py | 14 ++++++++------ .../tests/unit/edge3/cli/test_definition.py | 8 ++++---- .../edge3/tests/unit/edge3/cli/test_worker.py | 10 ++++++---- .../tests/unit/edge3/models/test_edge_worker.py | 16 ++++++++++------ 7 files changed, 34 insertions(+), 29 deletions(-) diff --git a/providers/edge3/docs/deployment.rst b/providers/edge3/docs/deployment.rst index d5e37d6041b4b..a48fea1bd53d7 100644 --- a/providers/edge3/docs/deployment.rst +++ b/providers/edge3/docs/deployment.rst @@ -250,8 +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, - ``-s``/``--state`` to filter by worker state, and ``-q``/``--queue`` to list only - workers serving a given queue. + ``-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 ad15af57ac049..b1b89bbb46949 100644 --- a/providers/edge3/src/airflow/providers/edge3/cli/definition.py +++ b/providers/edge3/src/airflow/providers/edge3/cli/definition.py @@ -57,11 +57,6 @@ metavar="PATTERN", help="Optional glob pattern to filter workers by name, e.g. 'prod-*'. Lists all workers if omitted.", ) -ARG_QUEUE = Arg( - ("-q", "--queue"), - metavar="QUEUE", - help="Optional queue name to filter workers by. Only workers serving this exact queue are listed.", -) ARG_REQUIRED_EDGE_HOSTNAME = Arg( ("-H", "--edge-hostname"), help="Set the hostname of worker if you have multiple workers on a single machine", @@ -194,7 +189,7 @@ ARG_OUTPUT, ARG_STATE, ARG_WORKER_NAME_PATTERN, - ARG_QUEUE, + 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 5bdb9dc82cbbd..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, queue=args.queue + 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 f36d662ce8216..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,7 +256,7 @@ def _fetch_edge_hosts_from_db( hostname: str | None = None, states: list | None = None, worker_name_pattern: str | None = None, - queue: str | None = None, + queues: list[str] | None = None, *, session: Session = NEW_SESSION, ) -> Sequence[EdgeWorkerModel]: @@ -271,10 +271,12 @@ def _fetch_edge_hosts_from_db( ) query = query.order_by(EdgeWorkerModel.worker_name) workers = session.scalars(query).all() - if queue: + 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. - workers = [worker for worker in workers if worker.queues and queue in worker.queues] + # 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 @@ -284,11 +286,11 @@ def get_registered_edge_hosts( *, states: list | None = None, worker_name_pattern: str | None = None, - queue: 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, queue=queue, session=session + states=states, worker_name_pattern=worker_name_pattern, queues=queues, session=session ) diff --git a/providers/edge3/tests/unit/edge3/cli/test_definition.py b/providers/edge3/tests/unit/edge3/cli/test_definition.py index 2cdaa38731341..cf6bb9e18ea7f 100644 --- a/providers/edge3/tests/unit/edge3/cli/test_definition.py +++ b/providers/edge3/tests/unit/edge3/cli/test_definition.py @@ -160,11 +160,11 @@ def test_list_workers_command_args(self): assert args.state == ["running", "maintenance"] assert args.worker_name_pattern == "prod-*" - def test_list_workers_command_queue_arg(self): - """Test list-workers command with the queue filter.""" - params = ["edge", "list-workers", "--queue", "gpu"] + 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.queue == "gpu" + assert args.queues == "gpu,default" def test_remote_edge_worker_request_maintenance_args(self): """Test remote-edge-worker-request-maintenance command with required arguments.""" diff --git a/providers/edge3/tests/unit/edge3/cli/test_worker.py b/providers/edge3/tests/unit/edge3/cli/test_worker.py index 5b075e29fb1d8..2a7c2675bb78e 100644 --- a/providers/edge3/tests/unit/edge3/cli/test_worker.py +++ b/providers/edge3/tests/unit/edge3/cli/test_worker.py @@ -1130,11 +1130,11 @@ 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-*", queue=None) + 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_queue(self, mock_edgeworker: EdgeWorkerModel): - args = self.parser.parse_args(["edge", "list-workers", "--output", "json", "--queue", "gpu"]) + 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( @@ -1146,7 +1146,9 @@ def test_list_edge_workers_passes_queue(self, mock_edgeworker: EdgeWorkerModel): ) as mock_get_hosts, ): edge_command.list_edge_workers(args) - mock_get_hosts.assert_called_once_with(states=None, worker_name_pattern=None, queue="gpu") + 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 e246d592d53f8..ab7c0550e7985 100644 --- a/providers/edge3/tests/unit/edge3/models/test_edge_worker.py +++ b/providers/edge3/tests/unit/edge3/models/test_edge_worker.py @@ -126,14 +126,18 @@ 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_queue_filters_by_exact_membership(self, session: Session): - hosts = get_registered_edge_hosts(queue="gpu", session=session) + 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_queue_no_match_returns_empty(self, session: Session): - hosts = get_registered_edge_hosts(queue="nonexistent", session=session) + 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_queue_combined_with_name_pattern(self, session: Session): - hosts = get_registered_edge_hosts(worker_name_pattern="prod-*", queue="gpu", session=session) + 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"}