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()

Reply via email to