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

kaxil 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 eebf253b72c Skip `on_kill()` for triggers reassigned to another 
triggerer (#73454)
eebf253b72c is described below

commit eebf253b72cc870f8725c10a6fcdd5f8058f2b9a
Author: Kaxil Naik <[email protected]>
AuthorDate: Wed Sep 23 22:53:21 2026 +0100

    Skip `on_kill()` for triggers reassigned to another triggerer (#73454)
    
    When a triggerer's heartbeat lapsed and its triggers were handed to another
    triggerer, the recovering triggerer treated every trigger missing from its
    assignment list as a user action and called on_kill(), cancelling remote 
work
    the new owner had just started polling. The supervisor now checks whether 
the
    trigger row still exists (and who owns it) before deciding: a surviving row 
is
    released locally without on_kill(); only a deleted row takes the user-action
    path.
---
 .../src/airflow/jobs/triggerer_job_runner.py       |  84 +++++++++---
 airflow-core/src/airflow/models/trigger.py         |   9 ++
 airflow-core/tests/unit/jobs/test_triggerer_job.py | 151 +++++++++++++++++++++
 airflow-core/tests/unit/models/test_trigger.py     |  17 +++
 4 files changed, 239 insertions(+), 22 deletions(-)

diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py 
b/airflow-core/src/airflow/jobs/triggerer_job_runner.py
index 876ca2e88b1..243b8f154d7 100644
--- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py
+++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py
@@ -150,6 +150,9 @@ logger = logging.getLogger(__name__)
 tracer = trace.get_tracer(__name__)
 # Private sentinel passed as the cancel message when a trigger is cancelled by 
user action
 _USER_ACTION_CANCEL_MSG = "__airflow_user_action__"
+# Private sentinel passed as the cancel message when a trigger's row moved to 
another triggerer, so
+# run_trigger() drops it locally without invoking on_kill()
+_REASSIGNED_CANCEL_MSG = "__airflow_reassigned__"
 
 _ON_CANCEL_TIMEOUT: int = conf.getint("triggerer", "on_kill_timeout", 
fallback=30)
 
@@ -336,7 +339,10 @@ class messages:
         type: Literal["TriggerStateSync"] = "TriggerStateSync"
 
         to_create: list[workloads.RunTrigger]
+        # Triggers whose row is gone (task no longer deferred on it); the 
runner invokes on_kill()
         to_cancel: set[int]
+        # Triggers whose row now belongs to another triggerer; the runner 
drops them without on_kill()
+        to_release: set[int] = Field(default_factory=set)
         # Seqs of shared-stream trigger events the supervisor has persisted
         # since the previous sync; the runner releases the matching broker
         # advances on receipt.
@@ -513,6 +519,8 @@ class TriggerRunnerSupervisor(WatchedSubprocess):
     # FinishedTriggers message
     cancelling_triggers: set[int] = attrs.field(factory=set, init=False)
 
+    releasing_triggers: set[int] = attrs.field(factory=set, init=False)
+
     # A list of RunTrigger workloads to send to the async process when it next 
checks in. We can't send it
     # directly as all comms has to be initiated by the subprocess
     creating_triggers: deque[workloads.RunTrigger] = 
attrs.field(factory=deque, init=False)
@@ -592,6 +600,7 @@ class TriggerRunnerSupervisor(WatchedSubprocess):
             for id in msg.finished or ():
                 self.running_triggers.discard(id)
                 self.cancelling_triggers.discard(id)
+                self.releasing_triggers.discard(id)
                 if factory := self.logger_cache.pop(id, None):
                     try:
                         factory.upload_to_remote()
@@ -610,6 +619,7 @@ class TriggerRunnerSupervisor(WatchedSubprocess):
             response = messages.TriggerStateSync(
                 to_create=[],
                 to_cancel=self.cancelling_triggers,
+                to_release=self.releasing_triggers,
                 events_persisted=events_persisted or None,
             )
 
@@ -923,6 +933,16 @@ class TriggerRunnerSupervisor(WatchedSubprocess):
         """Fetch trigger IDs associated with non-task entities."""
         return 
Trigger.fetch_trigger_ids_with_non_task_associations(session=session)
 
+    def fetch_trigger_assignments(self, trigger_ids: set[int]) -> dict[int, 
int | None]:
+        """
+        Map each trigger ID whose row still exists to the triggerer that owns 
it now.
+
+        Owns its own session so that ``update_triggers`` stays a pure 
diff/enqueue step that
+        subclasses without a metadata DB can keep calling after overriding 
this hook.
+        """
+        with create_session() as session:
+            return Trigger.fetch_assignments(trigger_ids, session=session)
+
     def build_trigger_workloads(self, new_trigger_ids: set[int]) -> 
list[workloads.RunTrigger]:
         """Build workloads for new trigger IDs."""
         dag_bag = DBDagBag()
@@ -977,12 +997,16 @@ class TriggerRunnerSupervisor(WatchedSubprocess):
         known_trigger_ids = self.running_triggers.union(
             (x[0] for x in self.events),
             self.cancelling_triggers,
+            self.releasing_triggers,
             (trigger[0] for trigger in self.failed_triggers),
             (trigger.id for trigger in self.creating_triggers),
         )
-        # Work out the two difference sets
+        # Work out the two difference sets. Triggers already queued for 
cancellation or release stay
+        # in running_triggers until the runner reports them finished; don't 
classify them again.
         new_trigger_ids = requested_trigger_ids - known_trigger_ids
-        cancel_trigger_ids = self.running_triggers - requested_trigger_ids
+        cancel_trigger_ids = (
+            self.running_triggers - requested_trigger_ids - 
self.cancelling_triggers - self.releasing_triggers
+        )
 
         if new_trigger_ids:
             workloads_to_create = self.build_trigger_workloads(new_trigger_ids)
@@ -995,8 +1019,17 @@ class TriggerRunnerSupervisor(WatchedSubprocess):
             self.creating_triggers.extend(workloads_to_create)
 
         if cancel_trigger_ids:
-            # Enqueue orphaned triggers for cancellation
-            self.cancelling_triggers.update(cancel_trigger_ids)
+            # Only the DB tells the two cases apart: a gone row means the task 
left the
+            # deferred state (user action), a surviving row means 
assign_unassigned handed
+            # the trigger to a triggerer that is now polling the remote work.
+            assignments = self.fetch_trigger_assignments(cancel_trigger_ids)
+            if assignments:
+                log.info(
+                    "Triggers were reassigned to another triggerer, releasing 
them without on_kill",
+                    new_owners=assignments,
+                )
+                self.releasing_triggers.update(assignments)
+            self.cancelling_triggers.update(cancel_trigger_ids - 
assignments.keys())
 
     def _register_pipe_readers(
         self,
@@ -1180,9 +1213,12 @@ class TriggerRunner:
     # Inbound queue of new triggers
     to_create: deque[workloads.RunTrigger]
 
-    # Inbound queue of deleted triggers
+    # Inbound queue of deleted triggers (user acted on the task; on_kill() 
runs)
     to_cancel: deque[int]
 
+    # Inbound queue of triggers reassigned to another triggerer (dropped 
locally; on_kill() skipped)
+    to_release: deque[int]
+
     # Outbound queue of events
     events: deque[TriggerEventEntry]
 
@@ -1207,6 +1243,7 @@ class TriggerRunner:
         self.trigger_cache = {}
         self.to_create = deque()
         self.to_cancel = deque()
+        self.to_release = deque()
         self.events = deque()
         self.failed_triggers = deque()
         self.team_name = None
@@ -1441,20 +1478,16 @@ class TriggerRunner:
             )
 
     async def cancel_triggers(self):
-        """
-        Drain the to_cancel queue and ensure all triggers that are not in the 
DB are cancelled.
-
-        This allows the cleanup job to delete them.
-        Passes "user-action" as the cancel message so that run_trigger() knows 
to invoke
-        on_kill(). Triggers in this queue are always removed because this is 
the path in which
-        the user performed some action on the task. Trigger redistribution 
goes through a separate
-        path.
-        """
-        while self.to_cancel:
-            trigger_id = self.to_cancel.popleft()
-            if trigger_id in self.triggers:
-                
self.triggers[trigger_id]["task"].cancel(_USER_ACTION_CANCEL_MSG)
-            await asyncio.sleep(0)
+        """Drain to_cancel and to_release; the cancel message tells 
run_trigger() if on_kill() applies."""
+        for queue, cancel_msg in (
+            (self.to_cancel, _USER_ACTION_CANCEL_MSG),
+            (self.to_release, _REASSIGNED_CANCEL_MSG),
+        ):
+            while queue:
+                trigger_id = queue.popleft()
+                if trigger_id in self.triggers:
+                    self.triggers[trigger_id]["task"].cancel(cancel_msg)
+                await asyncio.sleep(0)
 
     async def cleanup_finished_triggers(self) -> list[int]:
         """
@@ -1564,6 +1597,7 @@ class TriggerRunner:
         if resp:
             self.to_create.extend(resp.to_create)
             self.to_cancel.extend(resp.to_cancel)
+            self.to_release.extend(resp.to_release)
             if resp.events_persisted:
                 self._shared_streams.confirm_persisted(resp.events_persisted)
 
@@ -1692,11 +1726,13 @@ class TriggerRunner:
                     
self.events.append(TriggerEventEntry(trigger_id=trigger_id, event=event, 
persist_seq=seq))
                 span.set_status(Status(StatusCode.OK))
             except asyncio.CancelledError as e:
-                # A trigger can be cancelled for two reasons:
+                # A trigger can be cancelled for three reasons:
                 #   - The user acted on the task (mark failed / clear / mark 
succeeded).
+                #   - The trigger was reassigned to another triggerer (see
+                #     TriggerRunnerSupervisor.update_triggers); the new owner 
keeps running it.
                 #   - The triggerer is shutting down, here cancel_triggers() 
is not
                 #     involved — the shutdown path cancels tasks directly 
without a message.
-                # Only first case should invoke on_kill().
+                # Only the first case should invoke on_kill().
                 #
                 # For timeout, raise immediately without calling on_kill().
                 if timeout := timeout_after:
@@ -1705,7 +1741,11 @@ class TriggerRunner:
                         await self.log.aerror("Trigger cancelled due to 
timeout")
                         span.set_status(Status(StatusCode.ERROR), 
description=str(e))
                         raise
-                if e.args and e.args[0] == _USER_ACTION_CANCEL_MSG:
+                if e.args and e.args[0] == _REASSIGNED_CANCEL_MSG:
+                    await self.log.ainfo(
+                        "Trigger reassigned to another triggerer, dropping it 
without on_kill", name=name
+                    )
+                elif e.args and e.args[0] == _USER_ACTION_CANCEL_MSG:
                     await self.log.ainfo("Trigger cancelled by user action, 
invoking on_kill", name=name)
                     try:
                         await asyncio.wait_for(trigger.on_kill(), 
timeout=_ON_CANCEL_TIMEOUT)
diff --git a/airflow-core/src/airflow/models/trigger.py 
b/airflow-core/src/airflow/models/trigger.py
index 75578974db8..6a3e704fe79 100644
--- a/airflow-core/src/airflow/models/trigger.py
+++ b/airflow-core/src/airflow/models/trigger.py
@@ -368,6 +368,15 @@ class Trigger(Base):
 
         return list(session.scalars(query).all())
 
+    @classmethod
+    @provide_session
+    def fetch_assignments(
+        cls, ids: Iterable[int], *, session: Session = NEW_SESSION
+    ) -> dict[int, int | None]:
+        """Map each of ``ids`` that still has a trigger row to its current 
``triggerer_id``."""
+        rows = session.execute(select(cls.id, 
cls.triggerer_id).where(cls.id.in_(ids)))
+        return {trigger_id: triggerer_id for trigger_id, triggerer_id in rows}
+
     @classmethod
     @provide_session
     def assign_unassigned(
diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py 
b/airflow-core/tests/unit/jobs/test_triggerer_job.py
index 1d04c5c703b..79b1755cd4f 100644
--- a/airflow-core/tests/unit/jobs/test_triggerer_job.py
+++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py
@@ -46,6 +46,7 @@ from opentelemetry.sdk.trace.export import SimpleSpanProcessor
 from opentelemetry.sdk.trace.export.in_memory_span_exporter import 
InMemorySpanExporter
 from opentelemetry.trace.propagation.tracecontext import 
TraceContextTextMapPropagator
 from pydantic import TypeAdapter
+from sqlalchemy import delete, update
 from structlog.typing import FilteringBoundLogger
 
 from airflow._shared.timezones import timezone
@@ -54,6 +55,7 @@ from airflow.executors.workloads.task import TaskInstanceDTO
 from airflow.executors.workloads.trigger import RunTrigger
 from airflow.jobs.job import Job
 from airflow.jobs.triggerer_job_runner import (
+    _REASSIGNED_CANCEL_MSG,
     _USER_ACTION_CANCEL_MSG,
     ToTriggerRunner,
     ToTriggerSupervisor,
@@ -71,6 +73,7 @@ from airflow.models.dag_version import DagVersion
 from airflow.models.dagbag import DBDagBag
 from airflow.models.dagbundle import DagBundleModel
 from airflow.models.serialized_dag import SerializedDagModel
+from airflow.models.taskinstance import TaskInstance
 from airflow.models.xcom import XComModel
 from airflow.providers.standard.operators.empty import EmptyOperator
 from airflow.providers.standard.operators.python import PythonOperator
@@ -1476,6 +1479,63 @@ class TestTriggerRunner:
 
         mock_trigger.on_kill.assert_not_called()
 
+    def test_run_trigger_skips_on_kill_when_trigger_reassigned(self, session, 
cap_structlog) -> None:
+        """on_kill() is not called when the trigger was handed to another 
triggerer."""
+        trigger_runner = TriggerRunner()
+        trigger_runner.triggers = {
+            1: {"task": MagicMock(spec=asyncio.Task), "is_watcher": False, 
"name": "mock_name", "events": 0}
+        }
+        mock_trigger = MagicMock(spec=BaseTrigger)
+        mock_trigger.run.side_effect = 
asyncio.CancelledError(_REASSIGNED_CANCEL_MSG)
+        mock_trigger.task_instance = MagicMock()
+        mock_trigger.task_instance.map_index = -1
+        mock_trigger.on_kill = AsyncMock()
+
+        with pytest.raises(asyncio.CancelledError):
+            asyncio.run(trigger_runner.run_trigger(1, mock_trigger))
+
+        mock_trigger.on_kill.assert_not_called()
+        assert {
+            "event": "Trigger reassigned to another triggerer, dropping it 
without on_kill",
+            "log_level": "info",
+            "name": "mock_name",
+        } in cap_structlog
+
+    @pytest.mark.asyncio
+    async def 
test_cancel_triggers_uses_distinct_messages_for_user_action_and_reassignment(self)
 -> None:
+        trigger_runner = TriggerRunner()
+        user_task = MagicMock(spec=asyncio.Task)
+        reassigned_task = MagicMock(spec=asyncio.Task)
+        trigger_runner.triggers = {
+            1: {"task": user_task, "is_watcher": False, "name": "user", 
"events": 0},
+            2: {"task": reassigned_task, "is_watcher": False, "name": "moved", 
"events": 0},
+        }
+        trigger_runner.to_cancel.append(1)
+        trigger_runner.to_release.append(2)
+        # An id the runner no longer knows about is ignored on both queues.
+        trigger_runner.to_cancel.append(3)
+        trigger_runner.to_release.append(4)
+
+        await trigger_runner.cancel_triggers()
+
+        user_task.cancel.assert_called_once_with(_USER_ACTION_CANCEL_MSG)
+        reassigned_task.cancel.assert_called_once_with(_REASSIGNED_CANCEL_MSG)
+        assert not trigger_runner.to_cancel
+        assert not trigger_runner.to_release
+
+    @pytest.mark.asyncio
+    async def test_sync_state_to_supervisor_queues_released_triggers(self) -> 
None:
+        trigger_runner = TriggerRunner()
+        trigger_runner.comms_decoder = AsyncMock(spec=TriggerCommsDecoder)
+        trigger_runner.comms_decoder.asend.return_value = 
messages.TriggerStateSync(
+            to_create=[], to_cancel={1}, to_release={2}
+        )
+
+        await trigger_runner.sync_state_to_supervisor([])
+
+        assert list(trigger_runner.to_cancel) == [1]
+        assert list(trigger_runner.to_release) == [2]
+
     def 
test_run_trigger_on_kill_exception_does_not_swallow_cancelled_error(self, 
session) -> None:
         """CancelledError propagates even if on_kill() raises."""
         trigger_runner = TriggerRunner()
@@ -2779,6 +2839,97 @@ def 
test_update_triggers_skips_when_ti_has_no_dag_version(session, supervisor_bu
     supervisor.stdin.write.assert_not_called()
 
 
+def _alive_triggerer_job(session) -> Job:
+    other_job = Job(job_type="TriggererJob")
+    other_job.latest_heartbeat = timezone.utcnow()
+    session.add(other_job)
+    session.flush()
+    return other_job
+
+
+def 
test_load_triggers_releases_reassigned_trigger_without_user_action_cancel(session,
 supervisor_builder):
+    """
+    A trigger whose row moved to another triggerer must be dropped locally, 
not queued
+    as a user-action cancel (which would fire ``on_kill()`` and cancel remote 
work the new
+    owner is still polling).
+    """
+    trigger = TimeDeltaTrigger(datetime.timedelta(days=7))
+    _, _, trigger_orm, _ = create_trigger_in_db(session, trigger)
+    supervisor = supervisor_builder()
+    supervisor.running_triggers = {trigger_orm.id}
+
+    # Our heartbeat lapsed and another (alive) triggerer took the trigger over.
+    other_job = _alive_triggerer_job(session)
+    session.execute(update(Trigger).where(Trigger.id == 
trigger_orm.id).values(triggerer_id=other_job.id))
+    session.flush()
+
+    supervisor.load_triggers()
+
+    assert supervisor.cancelling_triggers == set()
+    assert supervisor.releasing_triggers == {trigger_orm.id}
+
+
+def test_load_triggers_cancels_deleted_trigger_as_user_action(session, 
supervisor_builder):
+    """A trigger whose row is gone (task left the deferred state) takes the 
user-action path."""
+    trigger = TimeDeltaTrigger(datetime.timedelta(days=7))
+    _, _, trigger_orm, task_instance = create_trigger_in_db(session, trigger)
+    supervisor = supervisor_builder()
+    supervisor.running_triggers = {trigger_orm.id}
+
+    session.execute(
+        update(TaskInstance)
+        .where(TaskInstance.id == task_instance.id)
+        .values(trigger_id=None, state="scheduled")
+    )
+    session.execute(delete(Trigger).where(Trigger.id == trigger_orm.id))
+    session.flush()
+
+    supervisor.load_triggers()
+
+    assert supervisor.cancelling_triggers == {trigger_orm.id}
+    assert supervisor.releasing_triggers == set()
+
+
+def 
test_update_triggers_splits_cancel_set_by_row_existence(supervisor_builder, 
mocker):
+    supervisor = supervisor_builder()
+    supervisor.running_triggers = {1, 2, 3}
+    fetch_assignments = mocker.patch.object(
+        TriggerRunnerSupervisor, "fetch_trigger_assignments", autospec=True, 
return_value={2: 99}
+    )
+
+    supervisor.update_triggers({3})
+
+    fetch_assignments.assert_called_once_with(supervisor, {1, 2})
+    assert supervisor.releasing_triggers == {2}
+    assert supervisor.cancelling_triggers == {1}
+
+    # Already-classified triggers are not re-queried on the next loop.
+    supervisor.update_triggers({3})
+    fetch_assignments.assert_called_once()
+
+
+def 
test_state_sync_sends_released_triggers_separately_from_cancels(supervisor_builder,
 mocker):
+    supervisor = supervisor_builder()
+    send_msg = mocker.patch.object(TriggerRunnerSupervisor, "send_msg", 
autospec=True)
+    log = MagicMock(spec=FilteringBoundLogger)
+    supervisor.running_triggers = {1, 2}
+    supervisor.cancelling_triggers = {1}
+    supervisor.releasing_triggers = {2}
+
+    supervisor._handle_request(messages.TriggerStateChanges(), log=log, 
req_id=1)
+
+    resp = send_msg.call_args.args[1]
+    assert isinstance(resp, messages.TriggerStateSync)
+    assert resp.to_cancel == {1}
+    assert resp.to_release == {2}
+
+    # The runner reporting them finished clears both tracking sets.
+    supervisor._handle_request(messages.TriggerStateChanges(finished=[1, 2]), 
log=log, req_id=2)
+    assert supervisor.cancelling_triggers == set()
+    assert supervisor.releasing_triggers == set()
+    assert supervisor.running_triggers == set()
+
+
 class TestTriggererJobRunner:
     @patch("airflow.jobs.triggerer_job_runner.stats.initialize")
     @patch.object(TriggerRunnerSupervisor, "start")
diff --git a/airflow-core/tests/unit/models/test_trigger.py 
b/airflow-core/tests/unit/models/test_trigger.py
index 0ab5a89e7ca..cf77cb39dd7 100644
--- a/airflow-core/tests/unit/models/test_trigger.py
+++ b/airflow-core/tests/unit/models/test_trigger.py
@@ -122,6 +122,23 @@ def 
test_fetch_trigger_ids_with_non_task_associations(session):
     assert results == {asset_trigger.id, callback_trigger.id}
 
 
+def test_fetch_assignments_maps_surviving_rows_to_their_triggerer(session):
+    owned = Trigger(classpath="airflow.triggers.testing.SuccessTrigger1", 
kwargs={})
+    owned.triggerer_id = 42
+    unassigned = Trigger(classpath="airflow.triggers.testing.SuccessTrigger2", 
kwargs={})
+    deleted = Trigger(classpath="airflow.triggers.testing.SuccessTrigger3", 
kwargs={})
+    session.add_all([owned, unassigned, deleted])
+    session.commit()
+    deleted_id = deleted.id
+    session.delete(deleted)
+    session.commit()
+
+    assignments = Trigger.fetch_assignments({owned.id, unassigned.id, 
deleted_id, deleted_id + 1000})
+
+    assert assignments == {owned.id: 42, unassigned.id: None}
+    assert Trigger.fetch_assignments(set()) == {}
+
+
 def test_clean_unused(session, dag_maker):
     """
     Tests that unused triggers (those with no task instances referencing them)

Reply via email to