This is an automated email from the ASF dual-hosted git repository.

ashb pushed a commit to branch local-executor-bookkeeping
in repository https://gitbox.apache.org/repos/asf/airflow.git

commit 59eb578e03e3d6713c1595392a96ee74d6e1458c
Author: Ash Berlin-Taylor <[email protected]>
AuthorDate: Sat Oct 3 09:50:24 2026 +0100

    Improve LocalExecutor bookkeeping: correctly add tasks to the `running` list
    
    Since #73916 landed, BaseExecutor now keeps a UUID-to-coordinates map for 
each
    attempt and drops an entry once the key is no longer queued, running or in 
the
    event buffer. The issue was that LocalExecutor never added dispatched work 
to
    `running`, so the entry was dropped while the task was still executing. The
    scheduler then received the final event with no coordinates ("Received 
executor
    event with state success for task instance <uuid> (coordinates=None)").
    
    This PR fixes that, and addresses a few gotchas that could, in edge cases
    around callbacks or post-task finalization, lead to a forever-dead/locked
    worker slot.
    
    Holding the key in `running` until the worker finishes needs LocalExecutor
    to know when that is, and workers share one activity queue so the parent
    cannot tell which worker took which workload. Results now carry the
    worker's pid and START tells the parent which pid owns the key. That lets
    it:
    
    - release the slot only on a terminal result from the owning worker,
    ignoring results from unknown pids or pids that don't own the key
    - fail the workload when its worker dies, instead of leaving it in
    `running` until the scheduler's heartbeat timeout
    
    A resumed task reuses its key, so a per-key dispatch count stops the old
    run's terminal result from releasing the new run's slot.
---
 .../src/airflow/executors/local_executor.py        |  74 +++-
 .../tests/unit/executors/test_local_executor.py    | 455 ++++++++++++++++++++-
 2 files changed, 510 insertions(+), 19 deletions(-)

diff --git a/airflow-core/src/airflow/executors/local_executor.py 
b/airflow-core/src/airflow/executors/local_executor.py
index 74c92a4f31a..a279073131f 100644
--- a/airflow-core/src/airflow/executors/local_executor.py
+++ b/airflow-core/src/airflow/executors/local_executor.py
@@ -38,6 +38,7 @@ import structlog
 
 from airflow.executors.base_executor import BaseExecutor, 
get_execution_api_server_url
 from airflow.executors.workloads import WorkloadType
+from airflow.executors.workloads.types import state_class_for_key
 
 # add logger to parameter of setproctitle to support logging
 if sys.platform == "darwin":
@@ -49,7 +50,10 @@ else:
 
 if TYPE_CHECKING:
     from airflow.executors.workloads import ExecutorWorkload
-    from airflow.executors.workloads.types import WorkloadResultType
+    from airflow.executors.workloads.types import WorkloadKey, WorkloadState
+    from airflow.models.taskinstance import TaskInstance
+
+    LocalResult = tuple[int, WorkloadKey, WorkloadState | None, Exception | 
None]
 
 
 def _get_executor_process_title_prefix(team_name: str | None) -> str:
@@ -65,7 +69,7 @@ def _get_executor_process_title_prefix(team_name: str | None) 
-> str:
 def _run_worker(
     logger_name: str,
     input: SimpleQueue[ExecutorWorkload | None],
-    output: Queue[WorkloadResultType],
+    output: Queue[LocalResult],
     unread_messages: multiprocessing.sharedctypes.Synchronized[int],
     team_conf,
 ):
@@ -98,8 +102,7 @@ def _run_worker(
             unread_messages.value -= 1
 
         key = LocalExecutor.get_workload_key(workload)
-        if workload.running_state is not None:
-            output.put((key, workload.running_state, None))
+        output.put((os.getpid(), key, workload.running_state, None))
 
         try:
             BaseExecutor.run_workload(
@@ -108,10 +111,10 @@ def _run_worker(
                 
proctitle=f"{_get_executor_process_title_prefix(team_conf.team_name)} 
{workload.display_name}",
                 subprocess_logs_to_stdout=True,
             )
-            output.put((key, workload.success_state, None))
+            output.put((os.getpid(), key, workload.success_state, None))
         except Exception as e:
             log.exception("Workload execution failed.", 
workload_type=type(workload).__name__)
-            output.put((key, workload.failure_state, e))
+            output.put((os.getpid(), key, workload.failure_state, e))
 
 
 class LocalExecutor(BaseExecutor):
@@ -136,12 +139,14 @@ class LocalExecutor(BaseExecutor):
     )
 
     activity_queue: SimpleQueue[ExecutorWorkload | None]
-    result_queue: SimpleQueue[WorkloadResultType]
+    result_queue: SimpleQueue[LocalResult]
     workers: dict[int, multiprocessing.Process]
     _unread_messages: multiprocessing.sharedctypes.Synchronized[int]
 
     def __init__(self, *args, **kwargs):
         super().__init__(*args, **kwargs)
+        self._worker_tasks: dict[int, WorkloadKey] = {}
+        self._dispatch_counts: dict[WorkloadKey, int] = {}
 
         # Resolve the start method at instantiation, not at import: the 
component CLI entry may have
         # set it via [<component>]/[core] mp_start_method before the executor 
is created.
@@ -165,6 +170,8 @@ class LocalExecutor(BaseExecutor):
         self.activity_queue = SimpleQueue()
         self.result_queue = SimpleQueue()
         self.workers = {}
+        self._worker_tasks.clear()
+        self._dispatch_counts.clear()
 
         # Mypy sees this value as `SynchronizedBase[c_uint]`, but that isn't 
the right runtime type behaviour
         # (it looks like an int to python)
@@ -177,10 +184,15 @@ class LocalExecutor(BaseExecutor):
             self._spawn_workers_with_gc_freeze(self.parallelism)
 
     def _check_workers(self):
+        self._read_results()
         # Reap any dead workers
         to_remove = set()
         for pid, proc in self.workers.items():
             if not proc.is_alive():
+                self._read_results()
+                # A worker killed between dequeue and START cannot identify 
its workload.
+                if (key := self._worker_tasks.pop(pid, None)) is not None:
+                    self._finish_dispatch(key, state_class_for_key(key).FAILED)
                 to_remove.add(pid)
                 proc.close()
 
@@ -246,14 +258,21 @@ class LocalExecutor(BaseExecutor):
 
     def sync(self) -> None:
         """Sync will get called periodically by the heartbeat method."""
-        self._read_results()
         self._check_workers()
 
     def _read_results(self):
         try:
             while not self.result_queue.empty():
-                key, state, exc = self.result_queue.get()
-                self.change_state(key, state)
+                pid, key, state, exc = self.result_queue.get()
+                if pid not in self.workers or key not in self.running:
+                    continue
+                if state is None or state == "running":
+                    self._worker_tasks[pid] = key
+                    if state is not None:
+                        self.change_state(key, state, remove_running=False)
+                elif self._worker_tasks.get(pid) == key:
+                    del self._worker_tasks[pid]
+                    self._finish_dispatch(key, state)
         except (OSError, EOFError):
             self.log.exception("Error reading from result queue")
 
@@ -316,11 +335,44 @@ class LocalExecutor(BaseExecutor):
 
     def _process_workloads(self, workload_list):
         for workload in workload_list:
-            self.activity_queue.put(workload)
             key = self.get_workload_key(workload)
+            self.activity_queue.put(workload)
             removed = self.executor_queues[workload.type].pop(key, None)
             if not removed:
                 raise KeyError(f"Workload {key} was not found in any queue")
+            self.running.add(key)
+            self._dispatch_counts[key] = self._dispatch_counts.get(key, 0) + 1
         with self._unread_messages:
             self._unread_messages.value += len(workload_list)
         self._check_workers()
+
+    def _finish_dispatch(self, key: WorkloadKey, state: WorkloadState) -> None:
+        # A resumed attempt reuses its key, so the previous dispatch can 
finish while the next one is live.
+        remaining = self._dispatch_counts.pop(key, 1) - 1
+        if remaining > 0:
+            self._dispatch_counts[key] = remaining
+        super().change_state(key, state, remove_running=remaining <= 0)
+
+    def _forget_workload(self, key: WorkloadKey) -> None:
+        self._dispatch_counts.pop(key, None)
+        self._worker_tasks = {
+            pid: task_key for pid, task_key in self._worker_tasks.items() if 
task_key != key
+        }
+
+    def change_state(self, key, state, info=None, remove_running=True) -> None:
+        if remove_running:
+            self._forget_workload(key)
+        super().change_state(key, state, info=info, 
remove_running=remove_running)
+
+    def fail_connection_test(self, key) -> None:
+        self._forget_workload(key)
+        super().fail_connection_test(key)
+
+    def revoke_task(self, *, ti: TaskInstance) -> None:
+        key = self.get_task_key(ti)
+        self.executor_queues[WorkloadType.EXECUTE_TASK].pop(key, None)
+        for pid, task_key in self._worker_tasks.items():
+            if task_key == key:
+                self._terminate_worker_process(self.workers[pid])
+        self._forget_workload(key)
+        self.running.discard(key)
diff --git a/airflow-core/tests/unit/executors/test_local_executor.py 
b/airflow-core/tests/unit/executors/test_local_executor.py
index b89a7841e16..b5213b3c640 100644
--- a/airflow-core/tests/unit/executors/test_local_executor.py
+++ b/airflow-core/tests/unit/executors/test_local_executor.py
@@ -20,6 +20,8 @@ from __future__ import annotations
 import gc
 import multiprocessing
 import os
+import signal
+import time
 from pathlib import Path
 from unittest import mock
 
@@ -27,10 +29,11 @@ import pytest
 from kgb import spy_on
 from uuid6 import uuid7
 
+import airflow.executors.local_executor as local_executor_module
 from airflow._shared.timezones import timezone
 from airflow.executors import workloads
 from airflow.executors.base_executor import BaseExecutor, ExecutorConf, 
get_execution_api_server_url
-from airflow.executors.local_executor import LocalExecutor
+from airflow.executors.local_executor import LocalExecutor, _run_worker
 from airflow.executors.workloads import WorkloadType
 from airflow.executors.workloads.base import BundleInfo
 from airflow.executors.workloads.callback import CallbackDTO
@@ -92,11 +95,65 @@ def _make_task_workload():
     )
 
 
-def _write_large_results_to_queue(result_queue, result_count, payload_size):
+def _write_large_results_to_queue(result_queue, activity_queue, 
unread_messages, result_count, payload_size):
     payload = RuntimeError("x" * payload_size)
     for _ in range(result_count):
-        key = uuid7()
-        result_queue.put((key, State.SUCCESS, payload))
+        workload = activity_queue.get()
+        with unread_messages:
+            unread_messages.value -= 1
+        key = LocalExecutor.get_workload_key(workload)
+        result_queue.put((os.getpid(), key, workload.running_state, None))
+        result_queue.put((os.getpid(), key, State.SUCCESS, payload))
+
+
+def _make_workload(kind):
+    if kind == "task":
+        return _make_task_workload()
+    if kind == "callback":
+        return workloads.ExecuteCallback(
+            callback=CallbackDTO(
+                id=uuid7(),
+                fetch_method=CallbackFetchMethod.IMPORT_PATH,
+                data={"path": "test.func", "kwargs": {}},
+            ),
+            dag_rel_path="test.py",
+            bundle_info=BundleInfo(name="bundle"),
+            token="token",
+            log_path=None,
+        )
+    return workloads.TestConnection(
+        connection_test_id=uuid7(), connection_id="test", timeout=10, 
token="token"
+    )
+
+
+def _hold_workload(workload, **kwargs):
+    Path(workload.token).touch()
+    signal.pause()
+
+
+def _run_blocking_worker(**kwargs):
+    with mock.patch.object(BaseExecutor, "run_workload", autospec=True, 
side_effect=_hold_workload):
+        _run_worker(**kwargs)
+
+
+def _add_mock_worker(executor, mocker, pid):
+    proc = mocker.create_autospec(multiprocessing.Process, instance=True)
+    proc.pid = pid
+    proc.is_alive.return_value = True
+    executor.workers[pid] = proc
+    return proc
+
+
[email protected]
+def local_executor_with_mock_worker(mocker):
+    mocker.patch.object(LocalExecutor, "_spawn_workers_with_gc_freeze", 
autospec=True)
+    mocker.patch.object(LocalExecutor, "_spawn_worker", autospec=True)
+    executor = LocalExecutor(parallelism=1)
+    executor.start()
+    proc = _add_mock_worker(executor, mocker, 12345)
+    yield executor, proc
+    executor.workers.clear()
+    executor.end()
 
 
 class TestLocalExecutor:
@@ -347,7 +404,7 @@ class TestLocalExecutor:
         assert proc.join.call_args_list == [mock.call(timeout=0.2), 
mock.call(timeout=0.2)]
 
     @pytest.mark.execution_timeout(10)
-    def test_end_drains_result_queue_to_avoid_join_deadlock(self):
+    def test_end_drains_result_queue_to_avoid_join_deadlock(self, mocker):
         # Pin the worker to "fork": the drain logic under test is 
start-method-agnostic, but under the
         # "forkserver" default (Python 3.14+ on Linux) each spawned worker 
re-imports the whole airflow
         # stack before it can write a result, which intermittently exceeds the 
execution_timeout and
@@ -355,13 +412,24 @@ class TestLocalExecutor:
         # immediately and reliably reproduces the full-result_queue scenario 
this test guards.
         ctx = multiprocessing.get_context("fork")
         executor = LocalExecutor(parallelism=1)
-        executor.activity_queue = ctx.SimpleQueue()
-        executor.result_queue = ctx.SimpleQueue()
+        mocker.patch.object(executor, "_spawn_workers_with_gc_freeze", 
autospec=True)
+        executor.start()
         result_count = 8
         payload_size = 128 * 1024
+        submitted = [_make_task_workload() for _ in range(result_count)]
+        for workload in submitted:
+            executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        with mock.patch.object(executor, "_check_workers", autospec=True):
+            executor._process_workloads(submitted)
         proc = ctx.Process(
             target=_write_large_results_to_queue,
-            args=(executor.result_queue, result_count, payload_size),
+            args=(
+                executor.result_queue,
+                executor.activity_queue,
+                executor._unread_messages,
+                result_count,
+                payload_size,
+            ),
         )
         proc.start()
         executor.workers = {proc.pid: proc}
@@ -369,6 +437,11 @@ class TestLocalExecutor:
         executor.end()
 
         assert len(executor.event_buffer) == result_count
+        assert set(executor.event_buffer) == 
{executor.get_task_key(workload.ti) for workload in submitted}
+        assert all(state == State.SUCCESS for state, _ in 
executor.event_buffer.values())
+        assert not executor.running
+        assert not executor._worker_tasks
+        assert executor._unread_messages.value == 0
 
     @pytest.mark.parametrize(
         ("conf_values", "expected_server"),
@@ -508,6 +581,372 @@ class TestLocalExecutor:
         executor.end()
 
 
+class TestLocalExecutorBookkeeping:
+    def test_dispatch_keeps_task_visible_without_a_worker_result(self, mocker):
+        mocker.patch.object(LocalExecutor, "_spawn_workers_with_gc_freeze", 
autospec=True)
+        mocker.patch.object(LocalExecutor, "_check_workers", autospec=True)
+        executor = LocalExecutor(parallelism=1)
+        executor.start()
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        try:
+            executor.heartbeat()
+            executor._drain_events_with_task_ids()
+
+            assert key in executor.running
+            assert executor.has_task(workload.ti)
+            assert executor.slots_available == 0
+            assert executor._task_coordinates[key] == workload.ti.key
+        finally:
+            executor.end()
+
+    def test_running_limits_later_heartbeats_and_reports_metrics(
+        self, local_executor_with_mock_worker, mocker
+    ):
+        executor, proc = local_executor_with_mock_worker
+        gauge = mocker.patch("airflow.executors.base_executor.stats.gauge", 
autospec=True)
+        first, second = _make_task_workload(), _make_task_workload()
+        executor.queue_workload(first, session=mock.create_autospec(Session, 
instance=True))
+        executor.heartbeat()
+        assert executor.slots_available == 0
+        executor.queue_workload(second, session=mock.create_autospec(Session, 
instance=True))
+
+        executor.heartbeat()
+
+        assert executor.running == {executor.get_task_key(first.ti)}
+        assert executor._unread_messages.value == 1
+        assert second in executor.executor_queues[second.type].values()
+        assert executor.has_task(first.ti)
+        metrics = {call.args[0]: call.kwargs["value"] for call in 
gauge.call_args_list[-3:]}
+        assert metrics == {"executor.open_slots": 0, "executor.queued_tasks": 
1, "executor.running_tasks": 1}
+
+    @pytest.mark.parametrize("kind", ["task", "callback", "connection"])
+    @pytest.mark.parametrize("succeeded", [True, False])
+    def test_start_retains_slot_and_terminal_clears_pid(
+        self, kind, succeeded, local_executor_with_mock_worker
+    ):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_workload(kind)
+        key = executor.get_workload_key(workload)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+
+        executor.result_queue.put((proc.pid, key, workload.running_state, 
None))
+        executor.sync()
+
+        assert executor._worker_tasks == {proc.pid: key}
+        assert key in executor.running
+        assert executor.slots_available == 0
+        if workload.running_state is None:
+            assert key not in executor.event_buffer
+        else:
+            assert executor.event_buffer[key] == (workload.running_state, None)
+        terminal = workload.success_state if succeeded else 
workload.failure_state
+        executor.result_queue.put((proc.pid, key, terminal, None))
+        executor.sync()
+        assert executor.event_buffer[key] == (terminal, None)
+        assert not executor._worker_tasks
+        assert not executor._dispatch_counts
+        assert executor.slots_available == 1
+
+    def test_result_uses_original_submitted_uuid_after_dto_changes(self, 
local_executor_with_mock_worker):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_task_workload()
+        key, coordinates = executor.get_task_key(workload.ti), workload.ti.key
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+        submitted = executor.activity_queue.get()
+        workload.ti.id = uuid7()
+        workload.ti.try_number += 1
+        assert executor.get_workload_key(submitted) == key
+        executor.result_queue.put((proc.pid, key, None, None))
+        executor.result_queue.put((proc.pid, key, workload.success_state, 
None))
+
+        executor.sync()
+        events, captured = executor._drain_events_with_task_ids()
+
+        assert events == {key: (workload.success_state, None)}
+        assert captured == {key: coordinates}
+        assert executor.slots_available == 1
+
+    def test_reaper_drains_start_sent_after_initial_poll(self, 
local_executor_with_mock_worker):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+        executor.activity_queue.get()
+        executor._unread_messages.value = 0
+
+        def died_after_start():
+            executor.result_queue.put((proc.pid, key, None, None))
+            return False
+
+        proc.is_alive.side_effect = died_after_start
+        executor.sync()
+
+        assert executor.event_buffer[key] == (workload.failure_state, None)
+        assert not executor.running
+        assert not executor._worker_tasks
+        proc.close.assert_called_once()
+
+    def test_revoke_task_releases_slot_of_workload_lost_before_start(self, 
local_executor_with_mock_worker):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+        executor.activity_queue.get()
+        executor._unread_messages.value = 0
+        proc.is_alive.return_value = False
+        executor.sync()
+        assert not executor.workers
+        assert key in executor.running
+
+        executor.revoke_task(ti=workload.ti)
+
+        assert not executor.running
+        assert not executor._dispatch_counts
+        assert executor.event_buffer == {}
+        assert executor.slots_available == 1
+
+    @pytest.mark.parametrize(
+        ("stage", "worker_terminated"),
+        [("queued", False), ("dispatched", False), ("started", True)],
+    )
+    def test_revoke_task_clears_workload_at_every_stage(
+        self, stage, worker_terminated, local_executor_with_mock_worker
+    ):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        if stage != "queued":
+            executor.heartbeat()
+        if stage == "started":
+            executor.result_queue.put((proc.pid, key, None, None))
+            executor.sync()
+            assert executor._worker_tasks == {proc.pid: key}
+
+        executor.revoke_task(ti=workload.ti)
+
+        assert not executor.executor_queues[workload.type]
+        assert not executor.running
+        assert not executor._worker_tasks
+        assert not executor._dispatch_counts
+        assert executor.event_buffer == {}
+        assert proc.terminate.called is worker_terminated
+
+    @pytest.mark.parametrize("kind", ["task", "connection"])
+    def test_external_timeout_clears_pid_and_rejects_late_results(
+        self, kind, local_executor_with_mock_worker
+    ):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_workload(kind)
+        key = executor.get_workload_key(workload)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+        executor.result_queue.put((proc.pid, key, workload.running_state, 
None))
+        executor.sync()
+
+        if kind == "connection":
+            executor.fail_connection_test(key)
+        else:
+            executor.change_state(key, workload.failure_state, 
remove_running=True)
+        executor.result_queue.put((proc.pid, key, workload.success_state, 
None))
+        executor.sync()
+
+        assert not executor._worker_tasks
+        assert executor.slots_available == 1
+        expected_state = workload.running_state if kind == "connection" else 
workload.failure_state
+        expected = {key: (expected_state, None)}
+        assert executor.event_buffer == expected
+        assert executor.workers[proc.pid] is proc
+
+    def test_one_worker_runs_workloads_back_to_back(self, 
local_executor_with_mock_worker):
+        executor, proc = local_executor_with_mock_worker
+        first, second = _make_task_workload(), _make_task_workload()
+        first_key, second_key = executor.get_task_key(first.ti), 
executor.get_task_key(second.ti)
+        for workload, key in ((first, first_key), (second, second_key)):
+            executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+            executor.heartbeat()
+            executor.result_queue.put((proc.pid, key, None, None))
+            executor.sync()
+            assert executor._worker_tasks == {proc.pid: key}
+            executor.result_queue.put((proc.pid, key, workload.success_state, 
None))
+            executor.sync()
+            assert not executor._worker_tasks
+        assert executor.event_buffer == {
+            first_key: (first.success_state, None),
+            second_key: (second.success_state, None),
+        }
+        assert executor.slots_available == 1
+
+    def test_redispatched_key_stays_tracked_after_previous_dispatch_finishes(
+        self, local_executor_with_mock_worker, mocker
+    ):
+        executor, first_proc = local_executor_with_mock_worker
+        second_proc = _add_mock_worker(executor, mocker, 54321)
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        session = mock.create_autospec(Session, instance=True)
+        executor.queue_workload(workload, session=session)
+        executor.heartbeat()
+        executor.queue_workload(workload, session=session)
+        executor._process_workloads([workload])
+        executor.result_queue.put((first_proc.pid, key, None, None))
+        executor.result_queue.put((first_proc.pid, key, 
workload.success_state, None))
+        executor.result_queue.put((second_proc.pid, key, None, None))
+
+        executor.sync()
+
+        assert executor.event_buffer[key] == (workload.success_state, None)
+        assert executor.has_task(workload.ti)
+        assert executor._worker_tasks == {second_proc.pid: key}
+        executor.result_queue.put((second_proc.pid, key, 
workload.failure_state, None))
+        executor.sync()
+        assert executor.event_buffer[key] == (workload.failure_state, None)
+        assert not executor.running
+        assert not executor._worker_tasks
+        assert not executor._dispatch_counts
+
+    def test_death_of_redispatched_workers_fails_key_after_last_dispatch(
+        self, local_executor_with_mock_worker, mocker
+    ):
+        executor, first_proc = local_executor_with_mock_worker
+        second_proc = _add_mock_worker(executor, mocker, 54321)
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        session = mock.create_autospec(Session, instance=True)
+        executor.queue_workload(workload, session=session)
+        executor.heartbeat()
+        executor.queue_workload(workload, session=session)
+        executor._process_workloads([workload])
+        executor.result_queue.put((first_proc.pid, key, None, None))
+        executor.result_queue.put((second_proc.pid, key, None, None))
+        executor.sync()
+        first_proc.is_alive.return_value = False
+
+        executor.sync()
+
+        assert key in executor.running
+        assert executor._worker_tasks == {second_proc.pid: key}
+        second_proc.is_alive.return_value = False
+        executor.sync()
+        assert executor.event_buffer[key] == (workload.failure_state, None)
+        assert not executor.running
+
+    def test_late_start_after_connection_test_reaped_is_ignored(self, 
local_executor_with_mock_worker):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_workload("connection")
+        key = executor.get_workload_key(workload)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+        executor.fail_connection_test(key)
+        executor.result_queue.put((proc.pid, key, workload.running_state, 
None))
+        executor.result_queue.put((proc.pid, key, workload.success_state, 
None))
+
+        executor.sync()
+
+        assert executor.event_buffer == {}
+        assert not executor._worker_tasks
+
+    def test_terminal_from_worker_that_does_not_own_the_key_is_ignored(
+        self, local_executor_with_mock_worker, mocker
+    ):
+        executor, owner = local_executor_with_mock_worker
+        other = _add_mock_worker(executor, mocker, 54321)
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+        executor.result_queue.put((owner.pid, key, None, None))
+        executor.result_queue.put((other.pid, key, workload.failure_state, 
None))
+
+        executor.sync()
+
+        assert executor.event_buffer == {}
+        assert executor._worker_tasks == {owner.pid: key}
+        assert key in executor.running
+
+    def test_result_from_unknown_pid_is_ignored(self, 
local_executor_with_mock_worker):
+        executor, proc = local_executor_with_mock_worker
+        workload = _make_task_workload()
+        key = executor.get_task_key(workload.ti)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        executor.heartbeat()
+        executor.result_queue.put((proc.pid + 1, key, None, None))
+
+        executor.sync()
+
+        assert not executor._worker_tasks
+        assert key in executor.running
+
+    def test_start_resets_bookkeeping_of_a_reused_executor(self, mocker):
+        mocker.patch.object(LocalExecutor, "_spawn_workers_with_gc_freeze", 
autospec=True)
+        executor = LocalExecutor(parallelism=1)
+        key = TaskInstanceUuid(uuid7())
+        executor._worker_tasks[12345] = key
+        executor._dispatch_counts[key] = 1
+
+        executor.start()
+
+        try:
+            assert not executor._worker_tasks
+            assert not executor._dispatch_counts
+        finally:
+            executor.end()
+
+    @pytest.mark.parametrize("start_method", ["fork", "spawn"])
+    @pytest.mark.parametrize("kind", ["task", "callback", "connection"])
+    @pytest.mark.execution_timeout(30)
+    def test_actual_worker_death_after_start_releases_slot(self, start_method, 
kind, mocker, tmp_path):
+        ctx = multiprocessing.get_context(start_method)
+        mocker.patch.object(
+            local_executor_module.multiprocessing,
+            "get_start_method",
+            autospec=True,
+            return_value=start_method,
+        )
+        mocker.patch.object(local_executor_module.multiprocessing, "Process", 
new=ctx.Process)
+        mocker.patch.object(local_executor_module.multiprocessing, "Value", 
new=ctx.Value)
+        mocker.patch.object(local_executor_module, "SimpleQueue", 
new=ctx.SimpleQueue)
+        mocker.patch.object(local_executor_module, "_run_worker", 
new=_run_blocking_worker)
+        executor = LocalExecutor(parallelism=1)
+        executor.start()
+        workload = _make_workload(kind)
+        marker = tmp_path / "entered"
+        workload.token = str(marker)
+        key = executor.get_workload_key(workload)
+        executor.queue_workload(workload, 
session=mock.create_autospec(Session, instance=True))
+        try:
+            executor.heartbeat()
+            deadline = time.monotonic() + 10
+            while not marker.exists():
+                assert time.monotonic() < deadline
+                executor.sync()
+                time.sleep(0.01)
+            executor.sync()
+            pid, proc = next(iter(executor.workers.items()))
+            assert executor._worker_tasks == {pid: key}
+            proc.kill()
+            proc.join(timeout=1)
+
+            executor.sync()
+
+            assert executor.event_buffer[key] == (workload.failure_state, None)
+            assert executor.slots_available == 1
+            assert not executor._worker_tasks
+            assert not executor.workers
+            executor.result_queue.put((pid, key, workload.success_state, None))
+            executor.sync()
+            assert executor.event_buffer[key] == (workload.failure_state, None)
+        finally:
+            executor.terminate()
+            executor.end()
+
+
 class TestLocalExecutorConnectionTestSupport:
     def test_test_connection_is_supported(self):
         executor = LocalExecutor()

Reply via email to