From 69af847148f82ce93f4674df12973a4bdc692eeb Mon Sep 17 00:00:00 2001 From: PoAn Yang Date: Tue, 25 Aug 2026 14:49:06 +0900 Subject: [PATCH 1/4] Make EdgeExecutor respect [core] parallelism Signed-off-by: PoAn Yang --- .../edge3/executors/edge_executor.py | 26 ++++++++++++-- .../providers/edge3/models/edge_job.py | 11 ++++-- .../edge3/executors/test_edge_executor.py | 35 +++++++++++++++++++ 3 files changed, 67 insertions(+), 5 deletions(-) diff --git a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py index e6f7af33eb06c..63e7d6547a31c 100644 --- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py +++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py @@ -28,6 +28,7 @@ from airflow.executors import workloads from airflow.executors.base_executor import BaseExecutor from airflow.models.taskinstance import TaskInstance +from airflow.models.taskinstancekey import TaskInstanceKey from airflow.providers.common.compat.sdk import Stats, timezone from airflow.providers.edge3.models.db import EdgeDBManager, check_db_manager_config from airflow.providers.edge3.models.edge_job import EdgeJobModel @@ -43,7 +44,6 @@ from sqlalchemy.orm import Session from airflow.cli.cli_config import GroupCommand - from airflow.models.taskinstancekey import TaskInstanceKey # TODO: Airflow 2 type hints; remove when Airflow 2 support is removed CommandType = Sequence[str] @@ -168,6 +168,7 @@ def queue_workload( team_name=self.team_name, ) ) + self.running.add(key) else: raise TypeError(f"Don't know how to queue workload of type {type(workload).__name__}") @@ -260,6 +261,26 @@ def _update_orphaned_jobs(self, session: Session) -> bool: return bool(lifeless_jobs) + def _get_tracked_job_keys(self, session: Session) -> set[TaskInstanceKey]: + """ + Read the keys of all jobs this team still has in the DB, queued ones included. + + Rows are read without locking on purpose: an edge worker fetches its next job with + ``FOR UPDATE SKIP LOCKED``, so locking the queued rows here would make it come back empty. + """ + return { + TaskInstanceKey(dag_id, task_id, run_id, try_number, map_index) + for dag_id, task_id, run_id, try_number, map_index in session.execute( + select( + EdgeJobModel.dag_id, + EdgeJobModel.task_id, + EdgeJobModel.run_id, + EdgeJobModel.try_number, + EdgeJobModel.map_index, + ).where(EdgeJobModel.team_name == self.team_name) + ) + } + def _purge_jobs(self, session: Session) -> bool: """Clean finished jobs.""" purged_marker = False @@ -284,8 +305,7 @@ def _purge_jobs(self, session: Session) -> bool: ).all() # Sync DB with executor otherwise runs out of sync in multi scheduler deployment - already_removed = self.running - set(job.key for job in jobs) - self.running = self.running - already_removed + self.running &= self._get_tracked_job_keys(session) for job in jobs: if job.key in self.running: diff --git a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py index 79576031112ea..202a7d31cf762 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py +++ b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py @@ -27,7 +27,8 @@ from sqlalchemy.orm import Mapped from airflow.models.base import StringID -from airflow.providers.common.compat.sdk import TaskInstanceKey, timezone +from airflow.models.taskinstancekey import TaskInstanceKey +from airflow.providers.common.compat.sdk import timezone from airflow.providers.common.compat.sqlalchemy.orm import mapped_column from airflow.providers.edge3.models.edge_base import Base from airflow.utils.log.logging_mixin import LoggingMixin @@ -92,7 +93,13 @@ def __init__( __table_args__ = (Index("rj_order", state, queued_dttm, queue),) @property - def key(self): + def key(self) -> TaskInstanceKey: + """ + Key of the job as the executor layer knows it. + + Deliberately the ``airflow.models`` class and not the ``airflow.sdk`` one: ``BaseExecutor`` + dispatches on it with ``isinstance``, and the two are unrelated classes. + """ return TaskInstanceKey(self.dag_id, self.task_id, self.run_id, self.try_number, self.map_index) @property diff --git a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py index 2f135957c4743..12deea06e9588 100644 --- a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py +++ b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py @@ -662,6 +662,41 @@ def test_queue_workload_execute_task(self): assert job.state == TaskInstanceState.QUEUED assert '"type":"ExecuteTask"' in job.command or '"type": "ExecuteTask"' in job.command + def test_queue_workload_occupies_an_executor_slot(self): + executor = EdgeExecutor() + workload = self._make_execute_task() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + assert workload.ti.key in executor.running + assert executor.slots_available == executor.parallelism - 1 + + # The slot stays taken while the job waits in the queue for a worker to pick it up. + executor.sync() + + assert workload.ti.key in executor.running + assert executor.slots_available == executor.parallelism - 1 + + @pytest.mark.parametrize("state", [TaskInstanceState.RUNNING, TaskInstanceState.SUCCESS]) + def test_sync_reports_state_of_queued_workload(self, state): + executor = EdgeExecutor() + workload = self._make_execute_task() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + executor.sync() + + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = state + session.commit() + + executor.sync() + + assert executor.get_event_buffer() == {workload.ti.key: (state, None)} + def test_queue_workload_execute_task_existing_job(self): executor = EdgeExecutor() workload = self._make_execute_task() From b9d711984a09f2e134c243117858d1b0766c0100 Mon Sep 17 00:00:00 2001 From: PoAn Yang Date: Tue, 8 Sep 2026 11:45:57 +0900 Subject: [PATCH 2/4] Update try_adopt_task_instances Signed-off-by: PoAn Yang --- .../edge3/executors/edge_executor.py | 25 +++++----- .../edge3/executors/test_edge_executor.py | 48 +++++++++++++++++++ 2 files changed, 62 insertions(+), 11 deletions(-) diff --git a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py index 63e7d6547a31c..b4c2a795a1bc4 100644 --- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py +++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py @@ -320,15 +320,13 @@ def _purge_jobs(self, session: Session) -> bool: if job.key in self.last_reported_state: del self.last_reported_state[job.key] self.success(job.key) - elif job.state in [ - TaskInstanceState.FAILED, - TaskInstanceState.RESTARTING, - TaskInstanceState.UP_FOR_RETRY, - ]: + elif job.state in [TaskInstanceState.FAILED, TaskInstanceState.UP_FOR_RETRY]: if job.key in self.last_reported_state: del self.last_reported_state[job.key] self.fail(job.key) else: + # RESTARTING is not a failure here: the fetch endpoint parks a claimed job in that + # state until the worker reports RUNNING. self.last_reported_state[job.key] = TaskInstanceState(job.state) if ( job.state == TaskInstanceState.SUCCESS @@ -405,17 +403,22 @@ def revoke_task(self, *, ti: TaskInstance, session: Session = NEW_SESSION): ) self.log.info("Revoked task instance %s from EdgeExecutor", ti.key) - def try_adopt_task_instances(self, tis: Sequence[TaskInstance]) -> Sequence[TaskInstance]: + @provide_session + def try_adopt_task_instances( + self, tis: Sequence[TaskInstance], *, session: Session = NEW_SESSION + ) -> Sequence[TaskInstance]: """ - Try to adopt running task instances that have been abandoned by a SchedulerJob dying. + Adopt the task instances whose job is still tracked in the edge_job table. - Anything that is not adopted will be cleared by the scheduler (and then become eligible for - re-scheduling) + The ``running`` set is empty after a scheduler restart, so the adopted keys go back into it + to keep slot accounting accurate. Task instances without a job row are returned so the + scheduler clears and re-schedules them. :return: any TaskInstances that were unable to be adopted """ - # We handle all running tasks from the DB in sync, no adoption logic needed. - return [] + tracked_keys = self._get_tracked_job_keys(session) + self.running.update(ti.key for ti in tis if ti.key in tracked_keys) + return [ti for ti in tis if ti.key not in tracked_keys] @staticmethod def get_cli_commands() -> list[GroupCommand]: diff --git a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py index 12deea06e9588..47601c68c767c 100644 --- a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py +++ b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py @@ -29,6 +29,7 @@ from sqlalchemy import delete, select from airflow.executors.workloads import BundleInfo, ExecuteTask +from airflow.models.taskinstance import TaskInstance from airflow.providers.common.compat.sdk import Stats, TaskInstanceKey, conf, timezone from airflow.providers.edge3.executors.edge_executor import EdgeExecutor from airflow.providers.edge3.models.edge_job import EdgeJobModel @@ -697,6 +698,53 @@ def test_sync_reports_state_of_queued_workload(self, state): assert executor.get_event_buffer() == {workload.ti.key: (state, None)} + def test_sync_keeps_slot_while_worker_claims_job(self): + executor = EdgeExecutor() + workload = self._make_execute_task() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + # The fetch endpoint parks a claimed job in RESTARTING until the worker reports RUNNING. + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = TaskInstanceState.RESTARTING + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {} + assert workload.ti.key in executor.running + assert executor.slots_available == executor.parallelism - 1 + + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = TaskInstanceState.RUNNING + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {workload.ti.key: (TaskInstanceState.RUNNING, None)} + + def test_try_adopt_task_instances_restores_slots_from_edge_job(self): + executor = EdgeExecutor() + workload = self._make_execute_task() + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + restarted_executor = EdgeExecutor() + tracked_ti = mock.Mock(spec=TaskInstance, key=workload.ti.key) + orphaned_ti = mock.Mock( + spec=TaskInstance, + key=TaskInstanceKey( + dag_id="test_dag", task_id="orphan", run_id="test_run", try_number=1, map_index=-1 + ), + ) + + not_adopted = restarted_executor.try_adopt_task_instances([tracked_ti, orphaned_ti]) + + assert not_adopted == [orphaned_ti] + assert restarted_executor.running == {workload.ti.key} + assert restarted_executor.slots_available == restarted_executor.parallelism - 1 + def test_queue_workload_execute_task_existing_job(self): executor = EdgeExecutor() workload = self._make_execute_task() From 679ec9e82318a4273a66bdf3eab47235e01f8af0 Mon Sep 17 00:00:00 2001 From: PoAn Yang Date: Wed, 16 Sep 2026 14:59:54 +0900 Subject: [PATCH 3/4] Count Edge callback workloads Signed-off-by: PoAn Yang --- .../edge3/executors/edge_executor.py | 51 +++++----- .../providers/edge3/models/edge_job.py | 34 +++++-- .../edge3/executors/test_edge_executor.py | 97 +++++++++---------- 3 files changed, 103 insertions(+), 79 deletions(-) diff --git a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py index b4c2a795a1bc4..971e3bc178f2a 100644 --- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py +++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py @@ -28,10 +28,9 @@ from airflow.executors import workloads from airflow.executors.base_executor import BaseExecutor from airflow.models.taskinstance import TaskInstance -from airflow.models.taskinstancekey import TaskInstanceKey from airflow.providers.common.compat.sdk import Stats, timezone from airflow.providers.edge3.models.db import EdgeDBManager, check_db_manager_config -from airflow.providers.edge3.models.edge_job import EdgeJobModel +from airflow.providers.edge3.models.edge_job import EdgeJobModel, build_job_key from airflow.providers.edge3.models.edge_logs import EdgeLogsModel from airflow.providers.edge3.models.edge_worker import EdgeWorkerModel, EdgeWorkerState, reset_metrics from airflow.providers.edge3.models.types import is_callback_execute @@ -44,6 +43,8 @@ from sqlalchemy.orm import Session from airflow.cli.cli_config import GroupCommand + from airflow.models.callback import CallbackKey + from airflow.models.taskinstancekey import TaskInstanceKey # TODO: Airflow 2 type hints; remove when Airflow 2 support is removed CommandType = Sequence[str] @@ -58,7 +59,7 @@ class EdgeExecutor(BaseExecutor): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.last_reported_state: dict[TaskInstanceKey, TaskInstanceState | str] = {} + self.last_reported_state: dict[TaskInstanceKey | CallbackKey, TaskInstanceState | str] = {} # Check if self has the ExecutorConf set on the self.conf attribute with all required methods. # In Airflow 2.x, ExecutorConf exists but lacks methods like getint, getboolean, getsection, etc. @@ -103,6 +104,7 @@ def queue_workload( session: Session, ) -> None: """Put new workload to queue. Airflow 3 entry point to execute a task.""" + key: TaskInstanceKey | CallbackKey if is_callback_execute(workload): from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG @@ -132,6 +134,7 @@ def queue_workload( team_name=self.team_name, ) ) + key = workload.key elif isinstance(workload, workloads.ExecuteTask): task_instance = workload.ti key = task_instance.key @@ -168,9 +171,10 @@ def queue_workload( team_name=self.team_name, ) ) - self.running.add(key) else: raise TypeError(f"Don't know how to queue workload of type {type(workload).__name__}") + # Added before the caller commits. On rollback, the reconciliation in _purge_jobs() drops the key. + self.running.add(key) def _process_workloads(self, workloads: Sequence[workloads.All]) -> None: """ @@ -261,25 +265,25 @@ def _update_orphaned_jobs(self, session: Session) -> bool: return bool(lifeless_jobs) - def _get_tracked_job_keys(self, session: Session) -> set[TaskInstanceKey]: + def _get_tracked_job_keys( + self, session: Session, states: Sequence[TaskInstanceState] | None = None + ) -> set[TaskInstanceKey | CallbackKey]: """ - Read the keys of all jobs this team still has in the DB, queued ones included. + Read the keys of this team's jobs still in the DB, optionally limited to those in ``states``. Rows are read without locking on purpose: an edge worker fetches its next job with ``FOR UPDATE SKIP LOCKED``, so locking the queued rows here would make it come back empty. """ - return { - TaskInstanceKey(dag_id, task_id, run_id, try_number, map_index) - for dag_id, task_id, run_id, try_number, map_index in session.execute( - select( - EdgeJobModel.dag_id, - EdgeJobModel.task_id, - EdgeJobModel.run_id, - EdgeJobModel.try_number, - EdgeJobModel.map_index, - ).where(EdgeJobModel.team_name == self.team_name) - ) - } + query = select( + EdgeJobModel.dag_id, + EdgeJobModel.task_id, + EdgeJobModel.run_id, + EdgeJobModel.try_number, + EdgeJobModel.map_index, + ).where(EdgeJobModel.team_name == self.team_name) + if states: + query = query.where(EdgeJobModel.state.in_(states)) + return {build_job_key(*row) for row in session.execute(query)} def _purge_jobs(self, session: Session) -> bool: """Clean finished jobs.""" @@ -408,15 +412,18 @@ def try_adopt_task_instances( self, tis: Sequence[TaskInstance], *, session: Session = NEW_SESSION ) -> Sequence[TaskInstance]: """ - Adopt the task instances whose job is still tracked in the edge_job table. + Adopt the task instances whose job is still in flight in the edge_job table. The ``running`` set is empty after a scheduler restart, so the adopted keys go back into it - to keep slot accounting accurate. Task instances without a job row are returned so the - scheduler clears and re-schedules them. + to keep slot accounting accurate. Task instances whose job is finished or missing are + returned so the scheduler clears and re-schedules them. :return: any TaskInstances that were unable to be adopted """ - tracked_keys = self._get_tracked_job_keys(session) + tracked_keys = self._get_tracked_job_keys( + session, + states=(TaskInstanceState.QUEUED, TaskInstanceState.RESTARTING, TaskInstanceState.RUNNING), + ) self.running.update(ti.key for ti in tis if ti.key in tracked_keys) return [ti for ti in tis if ti.key not in tracked_keys] diff --git a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py index 202a7d31cf762..073b26a6f08b5 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py +++ b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py @@ -17,6 +17,7 @@ from __future__ import annotations from datetime import datetime +from typing import TYPE_CHECKING from sqlalchemy import ( Index, @@ -31,9 +32,31 @@ from airflow.providers.common.compat.sdk import timezone from airflow.providers.common.compat.sqlalchemy.orm import mapped_column from airflow.providers.edge3.models.edge_base import Base +from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG +from airflow.providers.edge3.version_compat import AIRFLOW_V_3_3_PLUS from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.sqlalchemy import UtcDateTime +if TYPE_CHECKING: + from airflow.models.callback import CallbackKey + + +def build_job_key( + dag_id: str, task_id: str, run_id: str, try_number: int, map_index: int +) -> TaskInstanceKey | CallbackKey: + """ + Build the key the executor layer uses for a job row. + + A callback row carries the callback id in ``task_id``. A task row maps to the ``airflow.models`` + ``TaskInstanceKey``, not the ``airflow.sdk`` one, because ``BaseExecutor`` dispatches on it with + ``isinstance`` and the two are unrelated classes. + """ + if AIRFLOW_V_3_3_PLUS and dag_id == EXECUTE_CALLBACK_TAG: + from airflow.models.callback import CallbackKey + + return CallbackKey(id=task_id) + return TaskInstanceKey(dag_id, task_id, run_id, try_number, map_index) + class EdgeJobModel(Base, LoggingMixin): """ @@ -93,14 +116,9 @@ def __init__( __table_args__ = (Index("rj_order", state, queued_dttm, queue),) @property - def key(self) -> TaskInstanceKey: - """ - Key of the job as the executor layer knows it. - - Deliberately the ``airflow.models`` class and not the ``airflow.sdk`` one: ``BaseExecutor`` - dispatches on it with ``isinstance``, and the two are unrelated classes. - """ - return TaskInstanceKey(self.dag_id, self.task_id, self.run_id, self.try_number, self.map_index) + def key(self) -> TaskInstanceKey | CallbackKey: + """Key of the job as the executor layer knows it.""" + return build_job_key(self.dag_id, self.task_id, self.run_id, self.try_number, self.map_index) @property def last_update_t(self) -> float: diff --git a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py index 47601c68c767c..782c6179fae33 100644 --- a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py +++ b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py @@ -44,6 +44,7 @@ if AIRFLOW_V_3_3_PLUS: from airflow.executors.workloads import CallbackFetchMethod, ExecuteCallback, TaskInstanceDTO from airflow.executors.workloads.callback import CallbackDTO + from airflow.utils.state import CallbackState pytestmark = pytest.mark.db_test @@ -626,11 +627,11 @@ def setup(self): session.execute(delete(EdgeJobModel)) session.commit() - def _make_execute_task(self) -> ExecuteTask: + def _make_execute_task(self, task_id: str = "test_task") -> ExecuteTask: ti = TaskInstanceDTO( id=uuid4(), dag_version_id=uuid4(), - task_id="test_task", + task_id=task_id, dag_id="test_dag", run_id="test_run", try_number=1, @@ -647,6 +648,20 @@ def _make_execute_task(self) -> ExecuteTask: log_path="test.log", ) + def _make_execute_callback(self) -> ExecuteCallback: + callback = CallbackDTO( + id=str(uuid4()), + fetch_method=CallbackFetchMethod.IMPORT_PATH, + data={"path": "builtins.dict", "kwargs": {"a": 1, "b": 2, "c": 3}}, + ) + return ExecuteCallback( + callback=callback, + dag_rel_path=Path("test.py"), + bundle_info=BundleInfo(name="test_bundle", version="1.0"), + token="test_token", + log_path="test.log", + ) + def test_queue_workload_execute_task(self): executor = EdgeExecutor() workload = self._make_execute_task() @@ -663,27 +678,29 @@ def test_queue_workload_execute_task(self): assert job.state == TaskInstanceState.QUEUED assert '"type":"ExecuteTask"' in job.command or '"type": "ExecuteTask"' in job.command - def test_queue_workload_occupies_an_executor_slot(self): + @pytest.mark.parametrize("make_workload", ["_make_execute_task", "_make_execute_callback"]) + def test_queue_workload_occupies_an_executor_slot(self, make_workload): executor = EdgeExecutor() - workload = self._make_execute_task() + workload = getattr(self, make_workload)() with create_session() as session: executor.queue_workload(workload, session=session) session.commit() - assert workload.ti.key in executor.running + assert workload.key in executor.running assert executor.slots_available == executor.parallelism - 1 # The slot stays taken while the job waits in the queue for a worker to pick it up. executor.sync() - assert workload.ti.key in executor.running + assert workload.key in executor.running assert executor.slots_available == executor.parallelism - 1 + @pytest.mark.parametrize("make_workload", ["_make_execute_task", "_make_execute_callback"]) @pytest.mark.parametrize("state", [TaskInstanceState.RUNNING, TaskInstanceState.SUCCESS]) - def test_sync_reports_state_of_queued_workload(self, state): + def test_sync_reports_state_of_queued_workload(self, make_workload, state): executor = EdgeExecutor() - workload = self._make_execute_task() + workload = getattr(self, make_workload)() with create_session() as session: executor.queue_workload(workload, session=session) @@ -696,7 +713,8 @@ def test_sync_reports_state_of_queued_workload(self, state): executor.sync() - assert executor.get_event_buffer() == {workload.ti.key: (state, None)} + reported_states = TaskInstanceState if isinstance(workload, ExecuteTask) else CallbackState + assert executor.get_event_buffer() == {workload.key: (reported_states(state.value), None)} def test_sync_keeps_slot_while_worker_claims_job(self): executor = EdgeExecutor() @@ -723,15 +741,25 @@ def test_sync_keeps_slot_while_worker_claims_job(self): assert executor.get_event_buffer() == {workload.ti.key: (TaskInstanceState.RUNNING, None)} - def test_try_adopt_task_instances_restores_slots_from_edge_job(self): + @pytest.mark.parametrize( + "finished_state", [TaskInstanceState.SUCCESS, TaskInstanceState.FAILED, TaskInstanceState.REMOVED] + ) + def test_try_adopt_task_instances_restores_slots_from_edge_job(self, finished_state): executor = EdgeExecutor() - workload = self._make_execute_task() + queued = self._make_execute_task() + finished = self._make_execute_task(task_id="finished") with create_session() as session: - executor.queue_workload(workload, session=session) + executor.queue_workload(queued, session=session) + executor.queue_workload(finished, session=session) + session.commit() + with create_session() as session: + finished_job = session.scalar(select(EdgeJobModel).where(EdgeJobModel.task_id == "finished")) + finished_job.state = finished_state session.commit() restarted_executor = EdgeExecutor() - tracked_ti = mock.Mock(spec=TaskInstance, key=workload.ti.key) + queued_ti = mock.Mock(spec=TaskInstance, key=queued.key) + finished_ti = mock.Mock(spec=TaskInstance, key=finished.key) orphaned_ti = mock.Mock( spec=TaskInstance, key=TaskInstanceKey( @@ -739,10 +767,10 @@ def test_try_adopt_task_instances_restores_slots_from_edge_job(self): ), ) - not_adopted = restarted_executor.try_adopt_task_instances([tracked_ti, orphaned_ti]) + not_adopted = restarted_executor.try_adopt_task_instances([queued_ti, finished_ti, orphaned_ti]) - assert not_adopted == [orphaned_ti] - assert restarted_executor.running == {workload.ti.key} + assert not_adopted == [finished_ti, orphaned_ti] + assert restarted_executor.running == {queued.key} assert restarted_executor.slots_available == restarted_executor.parallelism - 1 def test_queue_workload_execute_task_existing_job(self): @@ -761,22 +789,7 @@ def test_queue_workload_execute_task_existing_job(self): def test_queue_workload_execute_callback(self): executor = EdgeExecutor() - id = str(uuid4()) - callback_data = CallbackDTO( - id=id, - fetch_method=CallbackFetchMethod.IMPORT_PATH, - data={ - "path": "builtins.dict", - "kwargs": {"a": 1, "b": 2, "c": 3}, - }, - ) - workload = ExecuteCallback( - callback=callback_data, - dag_rel_path=Path("test.py"), - bundle_info=BundleInfo(name="test_bundle", version="1.0"), - token="test_token", - log_path="test.log", - ) + workload = self._make_execute_callback() with create_session() as session: executor.queue_workload(workload, session=session) @@ -785,28 +798,14 @@ def test_queue_workload_execute_callback(self): job = session.scalar(select(EdgeJobModel)) assert job is not None assert job.dag_id == EXECUTE_CALLBACK_TAG - assert job.task_id == id - assert job.run_id == f"{EXECUTE_CALLBACK_TAG}-{id}" + assert job.task_id == workload.callback.id + assert job.run_id == f"{EXECUTE_CALLBACK_TAG}-{workload.callback.id}" assert job.state == TaskInstanceState.QUEUED assert '"type":"ExecuteCallback"' in job.command or '"type": "ExecuteCallback"' in job.command def test_queue_workload_execute_callback_existing_job(self): executor = EdgeExecutor() - callback_data = CallbackDTO( - id=str(uuid4()), - fetch_method=CallbackFetchMethod.IMPORT_PATH, - data={ - "path": "builtins.dict", - "kwargs": {"a": 1, "b": 2, "c": 3}, - }, - ) - workload = ExecuteCallback( - callback=callback_data, - dag_rel_path=Path("test.py"), - bundle_info=BundleInfo(name="test_bundle", version="1.0"), - token="test_token", - log_path="test.log", - ) + workload = self._make_execute_callback() with create_session() as session: executor.queue_workload(workload, session=session) From 9f61bbea4b8d169b4b9ada8a168c95b5c1009736 Mon Sep 17 00:00:00 2001 From: PoAn Yang Date: Thu, 17 Sep 2026 18:03:42 +0900 Subject: [PATCH 4/4] Identify Edge callback jobs by their full row identity Signed-off-by: PoAn Yang --- providers/edge3/docs/changelog.rst | 6 ++ .../edge3/executors/edge_executor.py | 56 ++++++++------ .../providers/edge3/models/edge_job.py | 10 +-- .../airflow/providers/edge3/models/types.py | 19 +++++ .../edge3/executors/test_edge_executor.py | 76 +++++++++++++++++-- .../tests/unit/edge3/models/test_edge_job.py | 29 ++++++- 6 files changed, 160 insertions(+), 36 deletions(-) diff --git a/providers/edge3/docs/changelog.rst b/providers/edge3/docs/changelog.rst index 5e513a51d58f6..990b2886eaf8d 100644 --- a/providers/edge3/docs/changelog.rst +++ b/providers/edge3/docs/changelog.rst @@ -27,6 +27,12 @@ Changelog --------- +.. warning:: + ``EdgeExecutor`` now counts the tasks and callbacks it has queued against ``[core] parallelism``, as the + other executors do. Until now that limit had no effect on Edge. If a scheduler keeps more than + ``parallelism`` (default 32) workloads in flight on Edge, raise ``[core] parallelism``. Otherwise the + scheduler leaves the rest in ``scheduled`` state until slots free up. + 4.3.2 ..... diff --git a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py index 971e3bc178f2a..0b5b9c9e4dd50 100644 --- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py +++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py @@ -33,7 +33,13 @@ from airflow.providers.edge3.models.edge_job import EdgeJobModel, build_job_key from airflow.providers.edge3.models.edge_logs import EdgeLogsModel from airflow.providers.edge3.models.edge_worker import EdgeWorkerModel, EdgeWorkerState, reset_metrics -from airflow.providers.edge3.models.types import is_callback_execute +from airflow.providers.edge3.models.types import ( + CALLBACK_JOB_MAP_INDEX, + CALLBACK_JOB_TRY_NUMBER, + EXECUTE_CALLBACK_TAG, + build_callback_run_id, + is_callback_execute, +) from airflow.utils.db import DBLocks, create_global_lock from airflow.utils.helpers import prune_dict from airflow.utils.session import NEW_SESSION, provide_session @@ -52,6 +58,17 @@ TaskTuple = tuple[TaskInstanceKey, CommandType, str | None, Any | None] +# _purge_jobs() reports on or deletes a job only while it is in one of these states. +_PURGE_HANDLED_STATES = ( + TaskInstanceState.RUNNING, + TaskInstanceState.SUCCESS, + TaskInstanceState.FAILED, + TaskInstanceState.REMOVED, + TaskInstanceState.RESTARTING, + TaskInstanceState.UP_FOR_RETRY, +) + + class EdgeExecutor(BaseExecutor): """Implementation of the EdgeExecutor to distribute work to Edge Workers via HTTP.""" @@ -106,13 +123,11 @@ def queue_workload( """Put new workload to queue. Airflow 3 entry point to execute a task.""" key: TaskInstanceKey | CallbackKey if is_callback_execute(workload): - from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG - existing_job = session.scalars( select(EdgeJobModel).where( EdgeJobModel.dag_id == EXECUTE_CALLBACK_TAG, EdgeJobModel.task_id == workload.callback.id, - EdgeJobModel.run_id == f"{EXECUTE_CALLBACK_TAG}-{workload.callback.id}", + EdgeJobModel.run_id == build_callback_run_id(workload.callback.id), ) ).first() @@ -124,9 +139,9 @@ def queue_workload( EdgeJobModel( dag_id=EXECUTE_CALLBACK_TAG, task_id=str(workload.callback.id), - run_id=f"{EXECUTE_CALLBACK_TAG}-{workload.callback.id}", - map_index=-1, - try_number=0, + run_id=build_callback_run_id(workload.callback.id), + map_index=CALLBACK_JOB_MAP_INDEX, + try_number=CALLBACK_JOB_TRY_NUMBER, queue=self.conf.get_mandatory_value("operators", "default_queue"), concurrency_slots=1, state=TaskInstanceState.QUEUED, @@ -266,10 +281,10 @@ def _update_orphaned_jobs(self, session: Session) -> bool: return bool(lifeless_jobs) def _get_tracked_job_keys( - self, session: Session, states: Sequence[TaskInstanceState] | None = None + self, session: Session, states: Sequence[TaskInstanceState] ) -> set[TaskInstanceKey | CallbackKey]: """ - Read the keys of this team's jobs still in the DB, optionally limited to those in ``states``. + Read the keys of this team's jobs that are in one of ``states``. Rows are read without locking on purpose: an edge worker fetches its next job with ``FOR UPDATE SKIP LOCKED``, so locking the queued rows here would make it come back empty. @@ -280,9 +295,7 @@ def _get_tracked_job_keys( EdgeJobModel.run_id, EdgeJobModel.try_number, EdgeJobModel.map_index, - ).where(EdgeJobModel.team_name == self.team_name) - if states: - query = query.where(EdgeJobModel.state.in_(states)) + ).where(EdgeJobModel.team_name == self.team_name, EdgeJobModel.state.in_(states)) return {build_job_key(*row) for row in session.execute(query)} def _purge_jobs(self, session: Session) -> bool: @@ -295,21 +308,16 @@ def _purge_jobs(self, session: Session) -> bool: .with_for_update(skip_locked=True) .where( EdgeJobModel.team_name == self.team_name, - EdgeJobModel.state.in_( - [ - TaskInstanceState.RUNNING, - TaskInstanceState.SUCCESS, - TaskInstanceState.FAILED, - TaskInstanceState.REMOVED, - TaskInstanceState.RESTARTING, - TaskInstanceState.UP_FOR_RETRY, - ] - ), + EdgeJobModel.state.in_(_PURGE_HANDLED_STATES), ) ).all() - # Sync DB with executor otherwise runs out of sync in multi scheduler deployment - self.running &= self._get_tracked_job_keys(session) + # Sync DB with executor otherwise runs out of sync in multi scheduler deployment. Only a queued job + # or one handled below keeps its slot. _update_orphaned_jobs() can leave a job in any task instance + # state, and a row this method never reads again would hold its slot until the scheduler restarts. + self.running &= self._get_tracked_job_keys( + session, states=(TaskInstanceState.QUEUED, *_PURGE_HANDLED_STATES) + ) for job in jobs: if job.key in self.running: diff --git a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py index 073b26a6f08b5..0309a57d6d434 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py +++ b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py @@ -32,7 +32,7 @@ from airflow.providers.common.compat.sdk import timezone from airflow.providers.common.compat.sqlalchemy.orm import mapped_column from airflow.providers.edge3.models.edge_base import Base -from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG +from airflow.providers.edge3.models.types import is_callback_job from airflow.providers.edge3.version_compat import AIRFLOW_V_3_3_PLUS from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.sqlalchemy import UtcDateTime @@ -47,11 +47,11 @@ def build_job_key( """ Build the key the executor layer uses for a job row. - A callback row carries the callback id in ``task_id``. A task row maps to the ``airflow.models`` - ``TaskInstanceKey``, not the ``airflow.sdk`` one, because ``BaseExecutor`` dispatches on it with - ``isinstance`` and the two are unrelated classes. + A row is a callback only if it has the full identity ``queue_workload()`` writes for callbacks, since + ``ExecuteCallback`` is a valid Dag id. A task row maps to the ``airflow.models`` ``TaskInstanceKey``, + not the ``airflow.sdk`` one, because ``BaseExecutor`` dispatches on it with ``isinstance``. """ - if AIRFLOW_V_3_3_PLUS and dag_id == EXECUTE_CALLBACK_TAG: + if AIRFLOW_V_3_3_PLUS and is_callback_job(dag_id, task_id, run_id, try_number, map_index): from airflow.models.callback import CallbackKey return CallbackKey(id=task_id) diff --git a/providers/edge3/src/airflow/providers/edge3/models/types.py b/providers/edge3/src/airflow/providers/edge3/models/types.py index 19cea39539dc9..dce93f52f827b 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/types.py +++ b/providers/edge3/src/airflow/providers/edge3/models/types.py @@ -44,3 +44,22 @@ def is_callback_execute(workload: workloads.All) -> TypeGuard[ExecuteCallback]: # This is the key used to identify execute_callback jobs. # Changing this value may break compatibility with existing data in the edge_job table. EXECUTE_CALLBACK_TAG = "ExecuteCallback" + +# The rest of the identity queue_workload() writes for a callback row. "ExecuteCallback" is a valid +# Dag id, so a row is a callback only when all four fields match. +CALLBACK_JOB_TRY_NUMBER = 0 +CALLBACK_JOB_MAP_INDEX = -1 + + +def build_callback_run_id(callback_id: str) -> str: + return f"{EXECUTE_CALLBACK_TAG}-{callback_id}" + + +def is_callback_job(dag_id: str, task_id: str, run_id: str, try_number: int, map_index: int) -> bool: + """Return whether a job row matches the identity ``queue_workload()`` writes for a callback.""" + return ( + dag_id == EXECUTE_CALLBACK_TAG + and run_id == build_callback_run_id(task_id) + and try_number == CALLBACK_JOB_TRY_NUMBER + and map_index == CALLBACK_JOB_MAP_INDEX + ) diff --git a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py index 782c6179fae33..fd08a503e554d 100644 --- a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py +++ b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py @@ -627,12 +627,12 @@ def setup(self): session.execute(delete(EdgeJobModel)) session.commit() - def _make_execute_task(self, task_id: str = "test_task") -> ExecuteTask: + def _make_execute_task(self, task_id: str = "test_task", dag_id: str = "test_dag") -> ExecuteTask: ti = TaskInstanceDTO( id=uuid4(), dag_version_id=uuid4(), task_id=task_id, - dag_id="test_dag", + dag_id=dag_id, run_id="test_run", try_number=1, map_index=-1, @@ -697,8 +697,16 @@ def test_queue_workload_occupies_an_executor_slot(self, make_workload): assert executor.slots_available == executor.parallelism - 1 @pytest.mark.parametrize("make_workload", ["_make_execute_task", "_make_execute_callback"]) - @pytest.mark.parametrize("state", [TaskInstanceState.RUNNING, TaskInstanceState.SUCCESS]) - def test_sync_reports_state_of_queued_workload(self, make_workload, state): + @pytest.mark.parametrize( + ("job_state", "reported_state"), + [ + (TaskInstanceState.RUNNING, "running"), + (TaskInstanceState.SUCCESS, "success"), + (TaskInstanceState.FAILED, "failed"), + (TaskInstanceState.UP_FOR_RETRY, "failed"), + ], + ) + def test_sync_reports_state_of_queued_workload(self, make_workload, job_state, reported_state): executor = EdgeExecutor() workload = getattr(self, make_workload)() @@ -708,13 +716,13 @@ def test_sync_reports_state_of_queued_workload(self, make_workload, state): executor.sync() with create_session() as session: - session.scalar(select(EdgeJobModel)).state = state + session.scalar(select(EdgeJobModel)).state = job_state session.commit() executor.sync() reported_states = TaskInstanceState if isinstance(workload, ExecuteTask) else CallbackState - assert executor.get_event_buffer() == {workload.key: (reported_states(state.value), None)} + assert executor.get_event_buffer() == {workload.key: (reported_states(reported_state), None)} def test_sync_keeps_slot_while_worker_claims_job(self): executor = EdgeExecutor() @@ -741,6 +749,62 @@ def test_sync_keeps_slot_while_worker_claims_job(self): assert executor.get_event_buffer() == {workload.ti.key: (TaskInstanceState.RUNNING, None)} + def test_sync_reports_job_that_finishes_after_being_marked_removed(self): + executor = EdgeExecutor() + workload = self._make_execute_callback() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + # When a callback job runs past the heartbeat timeout, _update_orphaned_jobs() marks it REMOVED. + for job_state in (TaskInstanceState.REMOVED, TaskInstanceState.SUCCESS): + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = job_state + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {workload.key: (CallbackState.SUCCESS, None)} + + @pytest.mark.parametrize( + "unhandled_state", + [TaskInstanceState.SCHEDULED, TaskInstanceState.DEFERRED, TaskInstanceState.UP_FOR_RESCHEDULE], + ) + def test_sync_frees_slot_of_job_in_state_purge_never_handles(self, unhandled_state): + executor = EdgeExecutor() + workload = self._make_execute_task() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + # _update_orphaned_jobs() copies the task instance state into the job, whatever that state is. + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = unhandled_state + session.commit() + executor.sync() + + assert workload.key not in executor.running + assert executor.slots_available == executor.parallelism + + def test_task_in_dag_named_after_callback_tag_keeps_its_task_key(self): + executor = EdgeExecutor() + workload = self._make_execute_task(dag_id=EXECUTE_CALLBACK_TAG) + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + executor.sync() + + assert workload.key in executor.running + + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = TaskInstanceState.RUNNING + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {workload.key: (TaskInstanceState.RUNNING, None)} + @pytest.mark.parametrize( "finished_state", [TaskInstanceState.SUCCESS, TaskInstanceState.FAILED, TaskInstanceState.REMOVED] ) diff --git a/providers/edge3/tests/unit/edge3/models/test_edge_job.py b/providers/edge3/tests/unit/edge3/models/test_edge_job.py index e0f75c25964c5..dc4b04135b98d 100644 --- a/providers/edge3/tests/unit/edge3/models/test_edge_job.py +++ b/providers/edge3/tests/unit/edge3/models/test_edge_job.py @@ -24,9 +24,12 @@ from sqlalchemy import delete, select from airflow.providers.common.compat.sdk import TaskInstanceKey -from airflow.providers.edge3.models.edge_job import EdgeJobModel +from airflow.providers.edge3.models.edge_job import EdgeJobModel, build_job_key +from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG from airflow.utils.state import TaskInstanceState +from tests_common.test_utils.version_compat import AIRFLOW_V_3_3_PLUS + if TYPE_CHECKING: from sqlalchemy.orm import Session @@ -63,6 +66,30 @@ def test_key_builds_task_instance_key(): assert job.key == TaskInstanceKey("test_dag", "test_task", "test_run", 2, 3) +@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="Callback workloads need Airflow 3.3+") +def test_build_job_key_maps_callback_row_to_callback_key(): + from airflow.models.callback import CallbackKey + + key = build_job_key(EXECUTE_CALLBACK_TAG, "abc", f"{EXECUTE_CALLBACK_TAG}-abc", 0, -1) + + assert key == CallbackKey(id="abc") + + +@pytest.mark.parametrize( + ("dag_id", "run_id", "try_number", "map_index"), + [ + pytest.param("test_dag", f"{EXECUTE_CALLBACK_TAG}-abc", 0, -1, id="other_dag"), + pytest.param(EXECUTE_CALLBACK_TAG, "manual__2026-01-01T00:00:00+00:00", 0, -1, id="task_run_id"), + pytest.param(EXECUTE_CALLBACK_TAG, f"{EXECUTE_CALLBACK_TAG}-abc", 1, -1, id="try_number"), + pytest.param(EXECUTE_CALLBACK_TAG, f"{EXECUTE_CALLBACK_TAG}-abc", 0, 2, id="map_index"), + ], +) +def test_build_job_key_keeps_task_key_unless_full_callback_identity(dag_id, run_id, try_number, map_index): + key = build_job_key(dag_id, "abc", run_id, try_number, map_index) + + assert key == TaskInstanceKey(dag_id, "abc", run_id, try_number, map_index) + + @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()