This is an automated email from the ASF dual-hosted git repository.
dheerajturaga pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 2aca1aa918e Add missing test modules for the edge3 provider (#73111)
2aca1aa918e is described below
commit 2aca1aa918ecb9465413a05b8b40a055a8c7a669
Author: Keith <[email protected]>
AuthorDate: Tue Sep 15 08:19:30 2026 +0900
Add missing test modules for the edge3 provider (#73111)
* Add missing test modules for the edge3 provider
The OVERLOOKED_TESTS guard list carries six edge3 entries. Five of the
test modules are genuinely missing; test_edge_worker.py already exists
on main but its exemption was never removed (the stale-entry check
compares Path objects against strings and never fires), so the module
methods and worker lifecycle operations it should cover were untested.
Batch the whole provider into one change, as requested by the
maintainers when closing the old tracking issue: add the five missing
test modules, extend test_edge_worker.py to cover the model queue
handling and the maintenance/shutdown/queue/concurrency operations, and
drop all six edge3 exemptions from the guard list.
* Make edge3 test factories explicit for mypy
The kwargs-dict factory pattern defeats mypy's heterogeneous-dict
inference and the Optional-returning lookup helper leaked union types
into assertions, both failing the providers type check in CI. Spell out
the factory parameters and split the worker lookup into an asserting
getter and an Optional finder.
---
.../tests/unit/always/test_project_structure.py | 6 -
.../edge3/cli/test_example_extended_sysinfo.py | 122 +++++++++++++++
.../edge3/tests/unit/edge3/models/test_edge_job.py | 130 ++++++++++++++++
.../tests/unit/edge3/models/test_edge_logs.py | 90 +++++++++++
.../tests/unit/edge3/models/test_edge_worker.py | 167 ++++++++++++++++++++-
.../tests/unit/edge3/worker_api/test_datamodels.py | 149 ++++++++++++++++++
.../unit/edge3/worker_api/test_datamodels_ui.py | 121 +++++++++++++++
7 files changed, 778 insertions(+), 7 deletions(-)
diff --git a/airflow-core/tests/unit/always/test_project_structure.py
b/airflow-core/tests/unit/always/test_project_structure.py
index 1e44ae13b85..2b6610e2e10 100644
--- a/airflow-core/tests/unit/always/test_project_structure.py
+++ b/airflow-core/tests/unit/always/test_project_structure.py
@@ -94,12 +94,6 @@ class TestProjectStructure:
"providers/common/compat/tests/unit/common/compat/standard/test_triggers.py",
"providers/common/compat/tests/unit/common/compat/standard/test_utils.py",
"providers/common/messaging/tests/unit/common/messaging/providers/test_sqs.py",
-
"providers/edge3/tests/unit/edge3/cli/test_example_extended_sysinfo.py",
- "providers/edge3/tests/unit/edge3/models/test_edge_job.py",
- "providers/edge3/tests/unit/edge3/models/test_edge_logs.py",
- "providers/edge3/tests/unit/edge3/models/test_edge_worker.py",
- "providers/edge3/tests/unit/edge3/worker_api/test_datamodels.py",
-
"providers/edge3/tests/unit/edge3/worker_api/test_datamodels_ui.py",
"providers/fab/tests/unit/fab/auth_manager/api_fastapi/datamodels/test_login.py",
"providers/fab/tests/unit/fab/migrations/test_env.py",
"providers/fab/tests/unit/fab/www/api_connexion/test_exceptions.py",
diff --git
a/providers/edge3/tests/unit/edge3/cli/test_example_extended_sysinfo.py
b/providers/edge3/tests/unit/edge3/cli/test_example_extended_sysinfo.py
new file mode 100644
index 00000000000..e4429b31bd6
--- /dev/null
+++ b/providers/edge3/tests/unit/edge3/cli/test_example_extended_sysinfo.py
@@ -0,0 +1,122 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+import logging
+import sys
+from unittest import mock
+
+import pytest
+
+from airflow.providers.edge3.cli.example_extended_sysinfo import
get_example_extended_sysinfo
+
+pytestmark = [pytest.mark.asyncio]
+
+MODULE = "airflow.providers.edge3.cli.example_extended_sysinfo"
+
+GIB = 1024**3
+
+
+class FakeAsyncPath:
+ """Stand-in for ``anyio.Path`` limited to what the sysinfo function
uses."""
+
+ def __init__(self, exists: bool = False, text: str = ""):
+ self._exists = exists
+ self._text = text
+
+ async def exists(self) -> bool:
+ return self._exists
+
+ async def read_text(self) -> str:
+ return self._text
+
+
+def _patch_system(cpu_usage: float, disk_free_gb: float, loadavg: float =
1.55):
+ return (
+ mock.patch(f"{MODULE}.shutil.disk_usage",
return_value=mock.Mock(free=disk_free_gb * GIB)),
+ mock.patch(f"{MODULE}.psutil.cpu_percent", return_value=cpu_usage),
+ mock.patch(f"{MODULE}.os.getloadavg", return_value=(loadavg, 1.0,
0.5)),
+ mock.patch(f"{MODULE}.Path", side_effect=lambda path:
FakeAsyncPath(exists=False)),
+ )
+
+
+async def test_reports_system_measurements():
+ patches = _patch_system(cpu_usage=10.0, disk_free_gb=100.0, loadavg=1.554)
+ with patches[0], patches[1], patches[2], patches[3]:
+ sysinfo = await get_example_extended_sysinfo()
+
+ assert sysinfo["platform"] == sys.platform
+ assert sysinfo["disk_free_gb"] == 100.0
+ assert sysinfo["cpu_usage"] == 10.0
+ assert sysinfo["sys_load"] == 1.55
+
+
[email protected](
+ ("cpu_usage", "disk_free_gb", "expected_status", "expected_text"),
+ [
+ (10.0, 100.0, logging.INFO, "I am good, sun is shining 🌞"),
+ (71.0, 100.0, logging.WARNING, "Warning condition!"),
+ (10.0, 19.0, logging.WARNING, "Warning condition!"),
+ (96.0, 100.0, logging.ERROR, "Critical condition!"),
+ (10.0, 4.0, logging.ERROR, "Critical condition!"),
+ ],
+)
+async def test_status_reflects_cpu_and_disk_thresholds(
+ cpu_usage, disk_free_gb, expected_status, expected_text
+):
+ patches = _patch_system(cpu_usage=cpu_usage, disk_free_gb=disk_free_gb)
+ with patches[0], patches[1], patches[2], patches[3]:
+ sysinfo = await get_example_extended_sysinfo()
+
+ assert sysinfo["status"] == expected_status
+ assert sysinfo["status_text"] == expected_text
+
+
+async def test_status_file_overrides_measured_status():
+ fake_paths = {
+ "/tmp/edge_error_status": FakeAsyncPath(exists=True, text="40"),
+ "/tmp/edge_error_status_text": FakeAsyncPath(exists=True, text="mocked
outage"),
+ }
+ patches = _patch_system(cpu_usage=10.0, disk_free_gb=100.0)
+ with (
+ patches[0],
+ patches[1],
+ patches[2],
+ mock.patch(f"{MODULE}.Path", side_effect=fake_paths.__getitem__),
+ ):
+ sysinfo = await get_example_extended_sysinfo()
+
+ assert sysinfo["status"] == 40
+ assert sysinfo["status_text"] == "mocked outage"
+
+
+async def test_status_file_without_text_file_drops_status_text():
+ fake_paths = {
+ "/tmp/edge_error_status": FakeAsyncPath(exists=True, text="40"),
+ "/tmp/edge_error_status_text": FakeAsyncPath(exists=False),
+ }
+ patches = _patch_system(cpu_usage=10.0, disk_free_gb=100.0)
+ with (
+ patches[0],
+ patches[1],
+ patches[2],
+ mock.patch(f"{MODULE}.Path", side_effect=fake_paths.__getitem__),
+ ):
+ sysinfo = await get_example_extended_sysinfo()
+
+ assert sysinfo["status"] == 40
+ assert "status_text" not in sysinfo
diff --git a/providers/edge3/tests/unit/edge3/models/test_edge_job.py
b/providers/edge3/tests/unit/edge3/models/test_edge_job.py
new file mode 100644
index 00000000000..e0f75c25964
--- /dev/null
+++ b/providers/edge3/tests/unit/edge3/models/test_edge_job.py
@@ -0,0 +1,130 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+from datetime import datetime, timezone as dt_timezone
+from typing import TYPE_CHECKING
+
+import pytest
+import time_machine
+from sqlalchemy import delete, select
+
+from airflow.providers.common.compat.sdk import TaskInstanceKey
+from airflow.providers.edge3.models.edge_job import EdgeJobModel
+from airflow.utils.state import TaskInstanceState
+
+if TYPE_CHECKING:
+ from sqlalchemy.orm import Session
+
+
+def _make_job(
+ *,
+ map_index: int = -1,
+ try_number: int = 1,
+ queued_dttm: datetime | None = None,
+ edge_worker: str | None = None,
+ last_update: datetime | None = None,
+ team_name: str | None = None,
+) -> EdgeJobModel:
+ return EdgeJobModel(
+ dag_id="test_dag",
+ task_id="test_task",
+ run_id="test_run",
+ map_index=map_index,
+ try_number=try_number,
+ state=TaskInstanceState.QUEUED,
+ queue="default",
+ concurrency_slots=1,
+ command="{}",
+ queued_dttm=queued_dttm,
+ edge_worker=edge_worker,
+ last_update=last_update,
+ team_name=team_name,
+ )
+
+
+def test_key_builds_task_instance_key():
+ job = _make_job(map_index=3, try_number=2)
+
+ assert job.key == TaskInstanceKey("test_dag", "test_task", "test_run", 2,
3)
+
+
+@time_machine.travel(datetime(2026, 1, 1, 12, 0, 0, tzinfo=dt_timezone.utc),
tick=False)
+def test_queued_dttm_defaults_to_now():
+ job = _make_job()
+
+ assert job.queued_dttm == datetime(2026, 1, 1, 12, 0, 0,
tzinfo=dt_timezone.utc)
+
+
+def test_queued_dttm_explicit_value_is_kept():
+ queued = datetime(2025, 6, 1, 8, 30, 0, tzinfo=dt_timezone.utc)
+
+ job = _make_job(queued_dttm=queued)
+
+ assert job.queued_dttm == queued
+
+
+def test_last_update_t_returns_timestamp_of_last_update():
+ last_update = datetime(2026, 1, 1, 12, 0, 0, tzinfo=dt_timezone.utc)
+
+ job = _make_job(last_update=last_update)
+
+ assert job.last_update_t == last_update.timestamp()
+
+
+@time_machine.travel(datetime(2026, 1, 1, 12, 0, 0, tzinfo=dt_timezone.utc),
tick=False)
+def test_last_update_t_falls_back_to_now_when_unset():
+ job = _make_job()
+
+ assert job.last_update_t == datetime.now().timestamp()
+
+
[email protected]_test
+class TestEdgeJobModelPersistence:
+ @pytest.fixture(autouse=True)
+ def _clean_table(self, session: Session):
+ session.execute(delete(EdgeJobModel))
+ session.commit()
+
+ def test_round_trip(self, session: Session):
+ queued = datetime(2026, 1, 1, 12, 0, 0, tzinfo=dt_timezone.utc)
+ session.add(
+ _make_job(
+ queued_dttm=queued,
+ edge_worker="worker-1",
+ team_name="team-a",
+ )
+ )
+ session.commit()
+ session.expunge_all()
+
+ job = session.scalars(select(EdgeJobModel)).one()
+ assert job.key == TaskInstanceKey("test_dag", "test_task", "test_run",
1, -1)
+ assert job.state == TaskInstanceState.QUEUED
+ assert job.queue == "default"
+ assert job.concurrency_slots == 1
+ assert job.queued_dttm == queued
+ assert job.edge_worker == "worker-1"
+ assert job.team_name == "team-a"
+
+ def test_try_numbers_are_separate_rows(self, session: Session):
+ session.add(_make_job(try_number=1))
+ session.add(_make_job(try_number=2))
+ session.commit()
+
+ try_numbers = set(session.scalars(select(EdgeJobModel.try_number)))
+ assert try_numbers == {1, 2}
diff --git a/providers/edge3/tests/unit/edge3/models/test_edge_logs.py
b/providers/edge3/tests/unit/edge3/models/test_edge_logs.py
new file mode 100644
index 00000000000..fc0801b421f
--- /dev/null
+++ b/providers/edge3/tests/unit/edge3/models/test_edge_logs.py
@@ -0,0 +1,90 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+from datetime import datetime, timedelta, timezone as dt_timezone
+from typing import TYPE_CHECKING
+
+import pytest
+from sqlalchemy import delete, select
+
+from airflow.providers.edge3.models.edge_logs import EdgeLogsModel
+
+if TYPE_CHECKING:
+ from sqlalchemy.orm import Session
+
+CHUNK_TIME = datetime(2026, 1, 1, 12, 0, 0, tzinfo=dt_timezone.utc)
+
+
+def _make_log_chunk(
+ *,
+ map_index: int = -1,
+ try_number: int = 1,
+ log_chunk_time: datetime = CHUNK_TIME,
+ log_chunk_data: str = "log line 1\n",
+) -> EdgeLogsModel:
+ return EdgeLogsModel(
+ dag_id="test_dag",
+ task_id="test_task",
+ run_id="test_run",
+ map_index=map_index,
+ try_number=try_number,
+ log_chunk_time=log_chunk_time,
+ log_chunk_data=log_chunk_data,
+ )
+
+
+def test_constructor_maps_all_fields():
+ chunk = _make_log_chunk(map_index=2, try_number=3)
+
+ assert chunk.dag_id == "test_dag"
+ assert chunk.task_id == "test_task"
+ assert chunk.run_id == "test_run"
+ assert chunk.map_index == 2
+ assert chunk.try_number == 3
+ assert chunk.log_chunk_time == CHUNK_TIME
+ assert chunk.log_chunk_data == "log line 1\n"
+
+
[email protected]_test
+class TestEdgeLogsModelPersistence:
+ @pytest.fixture(autouse=True)
+ def _clean_table(self, session: Session):
+ session.execute(delete(EdgeLogsModel))
+ session.commit()
+
+ def test_round_trip(self, session: Session):
+ session.add(_make_log_chunk())
+ session.commit()
+ session.expunge_all()
+
+ chunk = session.scalars(select(EdgeLogsModel)).one()
+ assert chunk.log_chunk_time == CHUNK_TIME
+ assert chunk.log_chunk_data == "log line 1\n"
+
+ def test_incremental_chunks_of_same_task_are_separate_rows(self, session:
Session):
+ session.add(_make_log_chunk())
+ session.add(
+ _make_log_chunk(
+ log_chunk_time=CHUNK_TIME + timedelta(seconds=10),
+ log_chunk_data="log line 2\n",
+ )
+ )
+ session.commit()
+
+ chunks =
session.scalars(select(EdgeLogsModel).order_by(EdgeLogsModel.log_chunk_time)).all()
+ assert [chunk.log_chunk_data for chunk in chunks] == ["log line 1\n",
"log line 2\n"]
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 ab7c0550e79..479255201ec 100644
--- a/providers/edge3/tests/unit/edge3/models/test_edge_worker.py
+++ b/providers/edge3/tests/unit/edge3/models/test_edge_worker.py
@@ -20,15 +20,23 @@ from typing import TYPE_CHECKING
from unittest import mock
import pytest
-from sqlalchemy import delete
+from sqlalchemy import delete, select
from airflow.providers.common.compat.sdk import Stats
from airflow.providers.edge3.models.edge_worker import (
EdgeWorkerModel,
EdgeWorkerState,
_glob_to_like_pattern,
+ add_worker_queues,
+ change_maintenance_comment,
+ exit_maintenance,
get_registered_edge_hosts,
+ remove_worker,
+ remove_worker_queues,
+ request_maintenance,
+ request_shutdown,
set_metrics,
+ set_worker_concurrency,
)
from tests_common.test_utils.version_compat import AIRFLOW_V_3_3_PLUS
@@ -141,3 +149,160 @@ class TestGetRegisteredEdgeHosts:
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"}
+
+
+class TestEdgeWorkerModelQueues:
+ def test_queues_default_to_none(self):
+ worker = EdgeWorkerModel(worker_name="worker-1",
state=EdgeWorkerState.IDLE, queues=None)
+
+ assert worker.queues is None
+
+ def test_queues_round_trip_through_string_column(self):
+ worker = EdgeWorkerModel(
+ worker_name="worker-1", state=EdgeWorkerState.IDLE,
queues=["default", "gpu"]
+ )
+
+ assert worker.queues == ["default", "gpu"]
+
+ def test_add_queues_deduplicates(self):
+ worker = EdgeWorkerModel(worker_name="worker-1",
state=EdgeWorkerState.IDLE, queues=["default"])
+
+ worker.add_queues(["gpu", "default"])
+
+ assert worker.queues is not None
+ assert sorted(worker.queues) == ["default", "gpu"]
+
+ def test_remove_queues_ignores_absent_queue(self):
+ worker = EdgeWorkerModel(
+ worker_name="worker-1", state=EdgeWorkerState.IDLE,
queues=["default", "gpu"]
+ )
+
+ worker.remove_queues(["gpu", "nonexistent"])
+
+ assert worker.queues == ["default"]
+
+ def test_update_state_coerces_string_to_enum(self):
+ worker = EdgeWorkerModel(worker_name="worker-1",
state=EdgeWorkerState.IDLE, queues=None)
+
+ worker.update_state("maintenance mode")
+
+ assert worker.state == EdgeWorkerState.MAINTENANCE_MODE
+
+
[email protected]_test
+class TestWorkerLifecycleOperations:
+ @pytest.fixture(autouse=True)
+ def setup_test_cases(self, session: Session):
+ session.execute(delete(EdgeWorkerModel))
+ session.add(
+ EdgeWorkerModel(worker_name="running-worker",
state=EdgeWorkerState.RUNNING, queues=["default"])
+ )
+ session.add(EdgeWorkerModel(worker_name="offline-worker",
state=EdgeWorkerState.OFFLINE, queues=None))
+ session.commit()
+
+ @staticmethod
+ def _find_worker(session: Session, worker_name: str) -> EdgeWorkerModel |
None:
+ return
session.scalar(select(EdgeWorkerModel).where(EdgeWorkerModel.worker_name ==
worker_name))
+
+ @classmethod
+ def _get_worker(cls, session: Session, worker_name: str) ->
EdgeWorkerModel:
+ worker = cls._find_worker(session, worker_name)
+ assert worker is not None
+ return worker
+
+ def test_request_maintenance_sets_state_and_comment(self, session:
Session):
+ request_maintenance("running-worker", "planned upgrade",
session=session)
+
+ worker = self._get_worker(session, "running-worker")
+ assert worker.state == EdgeWorkerState.MAINTENANCE_REQUEST
+ assert worker.maintenance_comment == "planned upgrade"
+
+ def test_exit_maintenance_sets_state_and_clears_comment(self, session:
Session):
+ request_maintenance("running-worker", "planned upgrade",
session=session)
+
+ exit_maintenance("running-worker", session=session)
+
+ worker = self._get_worker(session, "running-worker")
+ assert worker.state == EdgeWorkerState.MAINTENANCE_EXIT
+ assert worker.maintenance_comment is None
+
+ def test_change_maintenance_comment_in_maintenance_state(self, session:
Session):
+ request_maintenance("running-worker", "planned upgrade",
session=session)
+
+ change_maintenance_comment("running-worker", "upgrade extended",
session=session)
+
+ worker = self._get_worker(session, "running-worker")
+ assert worker.maintenance_comment == "upgrade extended"
+
+ def test_change_maintenance_comment_rejected_outside_maintenance(self,
session: Session):
+ with pytest.raises(TypeError, match="not in maintenance"):
+ change_maintenance_comment("running-worker", "some comment",
session=session)
+
+ def test_request_shutdown_sets_state(self, session: Session):
+ request_shutdown("running-worker", session=session)
+
+ worker = self._get_worker(session, "running-worker")
+ assert worker.state == EdgeWorkerState.SHUTDOWN_REQUEST
+
+ def test_request_shutdown_keeps_offline_worker_untouched(self, session:
Session):
+ request_shutdown("offline-worker", session=session)
+
+ worker = self._get_worker(session, "offline-worker")
+ assert worker.state == EdgeWorkerState.OFFLINE
+
+ def test_remove_worker_deletes_offline_worker(self, session: Session):
+ remove_worker("offline-worker", session=session)
+
+ assert self._find_worker(session, "offline-worker") is None
+
+ def test_remove_worker_rejects_active_worker(self, session: Session):
+ with pytest.raises(TypeError, match="Cannot remove edge worker"):
+ remove_worker("running-worker", session=session)
+
+ def test_add_worker_queues_extends_queues(self, session: Session):
+ add_worker_queues("running-worker", ["gpu"], session=session)
+
+ worker = self._get_worker(session, "running-worker")
+ assert worker.queues is not None
+ assert sorted(worker.queues) == ["default", "gpu"]
+
+ def test_remove_worker_queues_removes_queue(self, session: Session):
+ remove_worker_queues("running-worker", ["default"], session=session)
+
+ worker = self._get_worker(session, "running-worker")
+ assert worker.queues is None
+
+ def test_set_worker_concurrency_updates_value(self, session: Session):
+ set_worker_concurrency("running-worker", 8, session=session)
+
+ worker = self._get_worker(session, "running-worker")
+ assert worker.concurrency == 8
+
+ @pytest.mark.parametrize(
+ "operation",
+ [
+ lambda session: add_worker_queues("offline-worker", ["gpu"],
session=session),
+ lambda session: remove_worker_queues("offline-worker", ["gpu"],
session=session),
+ lambda session: set_worker_concurrency("offline-worker", 8,
session=session),
+ ],
+ )
+ def test_queue_and_concurrency_changes_rejected_for_offline_worker(self,
operation, session: Session):
+ with pytest.raises(TypeError):
+ operation(session)
+
+ @pytest.mark.parametrize(
+ "operation",
+ [
+ lambda session: request_maintenance("ghost", "comment",
session=session),
+ lambda session: exit_maintenance("ghost", session=session),
+ lambda session: change_maintenance_comment("ghost", "comment",
session=session),
+ lambda session: request_shutdown("ghost", session=session),
+ lambda session: remove_worker("ghost", session=session),
+ lambda session: add_worker_queues("ghost", ["gpu"],
session=session),
+ lambda session: remove_worker_queues("ghost", ["gpu"],
session=session),
+ lambda session: set_worker_concurrency("ghost", 8,
session=session),
+ ],
+ )
+ def test_unknown_worker_raises_value_error(self, operation, session:
Session):
+ with pytest.raises(ValueError, match="not found in list of registered
workers"):
+ operation(session)
diff --git a/providers/edge3/tests/unit/edge3/worker_api/test_datamodels.py
b/providers/edge3/tests/unit/edge3/worker_api/test_datamodels.py
new file mode 100644
index 00000000000..09321b8970d
--- /dev/null
+++ b/providers/edge3/tests/unit/edge3/worker_api/test_datamodels.py
@@ -0,0 +1,149 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+from datetime import datetime, timezone as dt_timezone
+
+import pytest
+from pydantic import ValidationError
+
+from airflow.providers.common.compat.sdk import TaskInstanceKey
+from airflow.providers.edge3.models.edge_worker import EdgeWorkerState
+from airflow.providers.edge3.worker_api.datamodels import (
+ EdgeJobBase,
+ EdgeJobFetched,
+ PushLogsBody,
+ WorkerQueuesBase,
+ WorkerRegistrationReturn,
+ WorkerSetStateReturn,
+ WorkerStateBody,
+)
+
+MOCK_COMMAND = {
+ "token": "mock",
+ "ti": {
+ "id": "4d828a62-a417-4936-a7a6-2b3fabacecab",
+ "task_id": "mock",
+ "dag_id": "mock",
+ "run_id": "mock",
+ "try_number": 1,
+ "dag_version_id": "01234567-89ab-cdef-0123-456789abcdef",
+ "pool_slots": 1,
+ "queue": "default",
+ "priority_weight": 1,
+ "start_date": "2023-01-01T00:00:00+00:00",
+ "map_index": -1,
+ },
+ "dag_rel_path": "mock.py",
+ "log_path": "mock.log",
+ "bundle_info": {"name": "hello", "version": "abc"},
+ "type": "ExecuteTask",
+}
+
+
+def _make_job_base(*, map_index: int = -1, try_number: int = 1) -> EdgeJobBase:
+ return EdgeJobBase(
+ dag_id="test_dag",
+ task_id="test_task",
+ run_id="test_run",
+ map_index=map_index,
+ try_number=try_number,
+ )
+
+
+def test_edge_job_base_key_builds_task_instance_key():
+ job = _make_job_base(map_index=3, try_number=2)
+
+ assert job.key == TaskInstanceKey("test_dag", "test_task", "test_run", 2,
3)
+
+
+class TestEdgeJobFetched:
+ def test_command_dict_is_coerced_to_workload(self):
+ job = EdgeJobFetched(
+ dag_id="test_dag",
+ task_id="test_task",
+ run_id="test_run",
+ map_index=-1,
+ try_number=1,
+ concurrency_slots=1,
+ command=MOCK_COMMAND, # type: ignore[arg-type]
+ )
+
+ assert job.command.ti.dag_id == "mock"
+
+ def test_identifier_names_all_key_components(self):
+ job = EdgeJobFetched(
+ dag_id="test_dag",
+ task_id="test_task",
+ run_id="test_run",
+ map_index=3,
+ try_number=2,
+ concurrency_slots=1,
+ command=MOCK_COMMAND, # type: ignore[arg-type]
+ )
+
+ assert job.identifier == (
+ "dag_id=test_dag task_id=test_task run_id=test_run map_index=3
try_number=2"
+ )
+
+
+def test_worker_queues_base_defaults():
+ body = WorkerQueuesBase()
+
+ assert body.queues is None
+ assert body.team_name is None
+
+
+class TestWorkerStateBody:
+ def test_defaults(self):
+ body = WorkerStateBody(state=EdgeWorkerState.IDLE, sysinfo={"status":
20})
+
+ assert body.jobs_active == 0
+ assert body.queues is None
+ assert body.maintenance_comments is None
+
+ def test_state_string_is_coerced_to_enum(self):
+ body = WorkerStateBody(state="maintenance mode", sysinfo={}) # type:
ignore[arg-type]
+
+ assert body.state == EdgeWorkerState.MAINTENANCE_MODE
+
+ def test_sysinfo_is_required(self):
+ with pytest.raises(ValidationError):
+ WorkerStateBody(state=EdgeWorkerState.IDLE) # type:
ignore[call-arg]
+
+
+def test_push_logs_body_parses_iso_timestamp():
+ body = PushLogsBody(
+ log_chunk_time="2026-01-01T12:00:00+00:00", # type: ignore[arg-type]
+ log_chunk_data="log line",
+ )
+
+ assert body.log_chunk_time == datetime(2026, 1, 1, 12, 0, 0,
tzinfo=dt_timezone.utc)
+
+
+def test_worker_registration_return_assumes_version_mismatch():
+ result = WorkerRegistrationReturn(last_update=datetime(2026, 1, 1,
tzinfo=dt_timezone.utc))
+
+ assert result.versions_match is False
+
+
+def test_worker_set_state_return_defaults():
+ result = WorkerSetStateReturn(state=EdgeWorkerState.IDLE, queues=None)
+
+ assert result.versions_match is False
+ assert result.maintenance_comments is None
+ assert result.concurrency is None
diff --git a/providers/edge3/tests/unit/edge3/worker_api/test_datamodels_ui.py
b/providers/edge3/tests/unit/edge3/worker_api/test_datamodels_ui.py
new file mode 100644
index 00000000000..017c34d9592
--- /dev/null
+++ b/providers/edge3/tests/unit/edge3/worker_api/test_datamodels_ui.py
@@ -0,0 +1,121 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+from datetime import datetime, timezone as dt_timezone
+
+import pytest
+from pydantic import ValidationError
+
+from airflow.providers.edge3.models.edge_worker import EdgeWorkerState
+from airflow.providers.edge3.worker_api.datamodels_ui import (
+ ConcurrencyRequest,
+ Job,
+ JobCollectionResponse,
+ MaintenanceRequest,
+ QueueUpdateRequest,
+ Worker,
+ WorkerCollectionResponse,
+)
+from airflow.utils.state import TaskInstanceState
+
+
+def _make_worker() -> Worker:
+ return Worker(worker_name="worker-1", state=EdgeWorkerState.IDLE,
sysinfo={"status": 20})
+
+
+def _make_job(*, queued_dttm: datetime | None = None, edge_worker: str | None
= None) -> Job:
+ return Job(
+ dag_id="test_dag",
+ task_id="test_task",
+ run_id="test_run",
+ map_index=-1,
+ try_number=1,
+ state=TaskInstanceState.RUNNING,
+ queue="default",
+ queued_dttm=queued_dttm,
+ edge_worker=edge_worker,
+ )
+
+
+class TestWorker:
+ def test_defaults(self):
+ worker = _make_worker()
+
+ assert worker.first_online is None
+ assert worker.last_heartbeat is None
+
+ def test_worker_name_is_required(self):
+ with pytest.raises(ValidationError):
+ Worker(state=EdgeWorkerState.IDLE, sysinfo={}) # type:
ignore[call-arg]
+
+
+class TestJob:
+ def test_defaults(self):
+ job = _make_job()
+
+ assert job.queued_dttm is None
+ assert job.edge_worker is None
+ assert job.last_update is None
+
+ def test_execution_fields_round_trip(self):
+ queued = datetime(2026, 1, 1, 12, 0, 0, tzinfo=dt_timezone.utc)
+
+ job = _make_job(queued_dttm=queued, edge_worker="worker-1")
+
+ assert job.state == TaskInstanceState.RUNNING
+ assert job.queue == "default"
+ assert job.queued_dttm == queued
+ assert job.edge_worker == "worker-1"
+
+
+def test_worker_collection_response_holds_workers():
+ response = WorkerCollectionResponse(workers=[_make_worker()],
total_entries=1)
+
+ assert response.workers[0].worker_name == "worker-1"
+ assert response.total_entries == 1
+
+
+def test_job_collection_response_holds_jobs():
+ response = JobCollectionResponse(jobs=[_make_job()], total_entries=1)
+
+ assert response.jobs[0].dag_id == "test_dag"
+ assert response.total_entries == 1
+
+
+def test_maintenance_request_requires_comment():
+ assert MaintenanceRequest(maintenance_comment="planned
upgrade").maintenance_comment == (
+ "planned upgrade"
+ )
+ with pytest.raises(ValidationError):
+ MaintenanceRequest() # type: ignore[call-arg]
+
+
+def test_queue_update_request_requires_queue_name():
+ assert QueueUpdateRequest(queue_name="gpu").queue_name == "gpu"
+ with pytest.raises(ValidationError):
+ QueueUpdateRequest() # type: ignore[call-arg]
+
+
+class TestConcurrencyRequest:
+ def test_positive_concurrency_is_accepted(self):
+ assert ConcurrencyRequest(concurrency=4).concurrency == 4
+
+ @pytest.mark.parametrize("concurrency", [0, -5])
+ def test_non_positive_concurrency_is_rejected(self, concurrency):
+ with pytest.raises(ValidationError):
+ ConcurrencyRequest(concurrency=concurrency)