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 656902875d5 Make EdgeExecutor respect [core] parallelism (#72048)
656902875d5 is described below
commit 656902875d59fae19c07ceb0a0e9b7eb0326c49a
Author: PoAn Yang <[email protected]>
AuthorDate: Sat Sep 19 05:58:46 2026 +0900
Make EdgeExecutor respect [core] parallelism (#72048)
* Make EdgeExecutor respect [core] parallelism
Signed-off-by: PoAn Yang <[email protected]>
* Update try_adopt_task_instances
Signed-off-by: PoAn Yang <[email protected]>
* Count Edge callback workloads
Signed-off-by: PoAn Yang <[email protected]>
* Identify Edge callback jobs by their full row identity
Signed-off-by: PoAn Yang <[email protected]>
---------
Signed-off-by: PoAn Yang <[email protected]>
---
providers/edge3/docs/changelog.rst | 6 +
.../providers/edge3/executors/edge_executor.py | 104 ++++++----
.../src/airflow/providers/edge3/models/edge_job.py | 31 ++-
.../src/airflow/providers/edge3/models/types.py | 19 ++
.../unit/edge3/executors/test_edge_executor.py | 218 +++++++++++++++++----
.../edge3/tests/unit/edge3/models/test_edge_job.py | 29 ++-
6 files changed, 334 insertions(+), 73 deletions(-)
diff --git a/providers/edge3/docs/changelog.rst
b/providers/edge3/docs/changelog.rst
index 5e513a51d58..990b2886eaf 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 e6f7af33eb0..0b5b9c9e4dd 100644
--- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py
+++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py
@@ -30,10 +30,16 @@ from airflow.executors.base_executor import BaseExecutor
from airflow.models.taskinstance import TaskInstance
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
+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
@@ -43,6 +49,7 @@ if TYPE_CHECKING:
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
@@ -51,6 +58,17 @@ if TYPE_CHECKING:
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."""
@@ -58,7 +76,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,14 +121,13 @@ class EdgeExecutor(BaseExecutor):
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
-
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()
@@ -122,9 +139,9 @@ class EdgeExecutor(BaseExecutor):
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,
@@ -132,6 +149,7 @@ class EdgeExecutor(BaseExecutor):
team_name=self.team_name,
)
)
+ key = workload.key
elif isinstance(workload, workloads.ExecuteTask):
task_instance = workload.ti
key = task_instance.key
@@ -170,6 +188,8 @@ class EdgeExecutor(BaseExecutor):
)
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:
"""
@@ -260,6 +280,24 @@ class EdgeExecutor(BaseExecutor):
return bool(lifeless_jobs)
+ def _get_tracked_job_keys(
+ self, session: Session, states: Sequence[TaskInstanceState]
+ ) -> set[TaskInstanceKey | CallbackKey]:
+ """
+ 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.
+ """
+ query = select(
+ EdgeJobModel.dag_id,
+ EdgeJobModel.task_id,
+ EdgeJobModel.run_id,
+ EdgeJobModel.try_number,
+ EdgeJobModel.map_index,
+ ).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:
"""Clean finished jobs."""
purged_marker = False
@@ -270,22 +308,16 @@ class EdgeExecutor(BaseExecutor):
.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
- already_removed = self.running - set(job.key for job in jobs)
- self.running = self.running - already_removed
+ # 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:
@@ -300,15 +332,13 @@ class EdgeExecutor(BaseExecutor):
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
@@ -385,17 +415,25 @@ class EdgeExecutor(BaseExecutor):
)
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 in flight 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 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
"""
- # We handle all running tasks from the DB in sync, no adoption logic
needed.
- return []
+ 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]
@staticmethod
def get_cli_commands() -> list[GroupCommand]:
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 79576031112..0309a57d6d4 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,
@@ -27,12 +28,35 @@ from sqlalchemy import (
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.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
+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 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 is_callback_job(dag_id, task_id, run_id,
try_number, map_index):
+ 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):
"""
@@ -92,8 +116,9 @@ class EdgeJobModel(Base, LoggingMixin):
__table_args__ = (Index("rj_order", state, queued_dttm, queue),)
@property
- def key(self):
- 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/src/airflow/providers/edge3/models/types.py
b/providers/edge3/src/airflow/providers/edge3/models/types.py
index 19cea39539d..dce93f52f82 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 2f135957c47..fd08a503e55 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 @@ import time_machine
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
@@ -43,6 +44,7 @@ from tests_common.test_utils.version_compat import
AIRFLOW_V_3_2_PLUS, AIRFLOW_V
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
@@ -625,12 +627,12 @@ class TestQueueWorkload:
session.execute(delete(EdgeJobModel))
session.commit()
- def _make_execute_task(self) -> 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="test_task",
- dag_id="test_dag",
+ task_id=task_id,
+ dag_id=dag_id,
run_id="test_run",
try_number=1,
map_index=-1,
@@ -646,6 +648,20 @@ class TestQueueWorkload:
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()
@@ -662,6 +678,165 @@ class TestQueueWorkload:
assert job.state == TaskInstanceState.QUEUED
assert '"type":"ExecuteTask"' in job.command or '"type":
"ExecuteTask"' in job.command
+ @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 = getattr(self, make_workload)()
+
+ with create_session() as session:
+ executor.queue_workload(workload, session=session)
+ session.commit()
+
+ 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.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(
+ ("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)()
+
+ 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 = job_state
+ session.commit()
+
+ executor.sync()
+
+ reported_states = TaskInstanceState if isinstance(workload,
ExecuteTask) else CallbackState
+ assert executor.get_event_buffer() == {workload.key:
(reported_states(reported_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_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]
+ )
+ def test_try_adopt_task_instances_restores_slots_from_edge_job(self,
finished_state):
+ executor = EdgeExecutor()
+ queued = self._make_execute_task()
+ finished = self._make_execute_task(task_id="finished")
+ with create_session() as 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()
+ 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(
+ 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([queued_ti,
finished_ti, orphaned_ti])
+
+ 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):
executor = EdgeExecutor()
workload = self._make_execute_task()
@@ -678,22 +853,7 @@ class TestQueueWorkload:
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)
@@ -702,28 +862,14 @@ class TestQueueWorkload:
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)
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 e0f75c25964..dc4b04135b9 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 @@ 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.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)
[email protected](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")
+
+
[email protected](
+ ("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()