From b18320b1e7593f9e4d0032915ba280c5cb11127a Mon Sep 17 00:00:00 2001 From: Idan Izhaki Date: Mon, 13 Jul 2026 17:55:52 -0700 Subject: [PATCH 1/5] Keep triggerer running when deferred tasks disappear A stale task reference in a pinned serialized Dag can otherwise abort workload construction for every trigger in the supervisor batch and cause repeated triggerer restarts. --- .../src/airflow/jobs/triggerer_job_runner.py | 19 ++- .../tests/unit/jobs/test_triggerer_job.py | 144 ++++++++++++++++++ 2 files changed, 161 insertions(+), 2 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 355763aed5f13..27318adab63ad 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -50,6 +50,7 @@ from airflow._shared.observability.metrics import stats from airflow._shared.timezones import timezone from airflow.configuration import conf +from airflow.exceptions import TaskNotFound from airflow.executors import workloads from airflow.executors.workloads.task import TaskInstanceDTO from airflow.jobs.base_job_runner import BaseJobRunner @@ -879,12 +880,26 @@ def _create_workload( ) if serialized_dag_model: - task = serialized_dag_model.dag.get_task(trigger.task_instance.task_id) + try: + task = serialized_dag_model.dag.get_task(trigger.task_instance.task_id) + except TaskNotFound: + # Avoid this lookup for ordinary deferrals once trigger context requirements are persisted; + # tracked at https://github.com/apache/airflow/issues/69841 + log.warning( + "Task for deferred trigger was not found in serialized Dag; " + "starting trigger without Dag context", + trigger_id=trigger.id, + ti_id=trigger.task_instance.id, + dag_id=trigger.task_instance.dag_id, + task_id=trigger.task_instance.task_id, + dag_version_id=trigger.task_instance.dag_version_id, + ) + task = None # When a TaskInstance of a Trigger contains a task with start_from_trigger enabled, # it means we need to load the SerializedDagModel so we can build a RuntimeTaskInstance later on which # will allow us to build a context on which we will render the templated fields. - if task.start_from_trigger: + if task is not None and task.start_from_trigger: log.info("Start from trigger enabled for task %s", task.task_id) dag_run = trigger.task_instance.get_dagrun(session=session) diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 0f307fc150520..3866c563e42a0 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -30,6 +30,7 @@ import uuid from collections.abc import AsyncIterator from socket import socket, socketpair +from types import SimpleNamespace from typing import TYPE_CHECKING, Any from unittest import mock from unittest.mock import ANY, AsyncMock, MagicMock, patch @@ -49,6 +50,7 @@ from structlog.typing import FilteringBoundLogger from airflow._shared.timezones import timezone +from airflow.exceptions import TaskNotFound from airflow.executors import workloads from airflow.executors.workloads.task import TaskInstanceDTO from airflow.executors.workloads.trigger import RunTrigger @@ -68,6 +70,7 @@ ) from airflow.models import Connection, DagModel, DagRun, Trigger, Variable 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.xcom import XComModel @@ -609,6 +612,147 @@ def test_create_workload_uses_supervisor_id_without_job(jobless_supervisor, mock assert factory.log_path == f"/logs/ti.trigger.{jobless_supervisor.id}.log" +def test_create_workload_missing_task_returns_context_free_workload(jobless_supervisor, mocker, caplog): + ti = SimpleNamespace( + id=uuid.uuid4(), + dag_id="example_dag", + dag_version_id=uuid.uuid4(), + task_id="removed_task", + trigger_timeout=None, + ) + trigger = SimpleNamespace( + id=7, + classpath="some.path.Trigger", + encrypted_kwargs="encrypted", + task_instance=ti, + ) + + ser_ti = mocker.Mock(spec=TaskInstanceDTO) + mocker.patch( + "airflow.jobs.triggerer_job_runner.TaskInstanceDTO.model_validate", + autospec=True, + return_value=ser_ti, + ) + + def get_task(task_id): + raise TaskNotFound(task_id) + + serialized_dag_model = SimpleNamespace(dag=SimpleNamespace(get_task=get_task)) + dag_bag = mocker.Mock(spec=DBDagBag) + dag_bag.get_serialized_dag_model.return_value = serialized_dag_model + + workload = jobless_supervisor._create_workload( + trigger=trigger, + dag_bag=dag_bag, + render_log_fname=lambda *, ti: "/logs/ti", + session=mocker.sentinel.session, + ) + + assert workload == RunTrigger( + id=trigger.id, + classpath=trigger.classpath, + encrypted_kwargs=trigger.encrypted_kwargs, + ti=ser_ti, + ) + assert { + "event": "Task for deferred trigger was not found in serialized Dag; " + "starting trigger without Dag context", + "trigger_id": trigger.id, + "ti_id": trigger.task_instance.id, + "dag_id": trigger.task_instance.dag_id, + "task_id": trigger.task_instance.task_id, + "dag_version_id": trigger.task_instance.dag_version_id, + } in caplog + + +def test_build_trigger_workloads_continues_after_missing_task(jobless_supervisor, mocker): + def get_missing_task(task_id): + raise TaskNotFound(task_id) + + stale_ti = SimpleNamespace( + id=uuid.uuid4(), + dag_id="example_dag", + dag_version_id=uuid.uuid4(), + task_id="removed_task", + trigger_timeout=None, + ) + stale_trigger = SimpleNamespace( + id=1, + classpath="some.path.StaleTrigger", + encrypted_kwargs="stale", + task_instance=stale_ti, + ) + healthy_ti = SimpleNamespace( + id=uuid.uuid4(), + dag_id="example_dag", + dag_version_id=uuid.uuid4(), + task_id="healthy_task", + trigger_timeout=None, + get_dagrun=lambda *, session: SimpleNamespace( + dag_run_data=SimpleNamespace( + model_dump=lambda **kwargs: {"dag_id": "example_dag"}, + ) + ), + ) + healthy_trigger = SimpleNamespace( + id=2, + classpath="some.path.HealthyTrigger", + encrypted_kwargs="healthy", + task_instance=healthy_ti, + ) + + serialized_stale_dag = SimpleNamespace(dag=SimpleNamespace(get_task=get_missing_task)) + serialized_healthy_dag = SimpleNamespace( + data={"dag": "serialized"}, + dag=SimpleNamespace( + get_task=lambda task_id: SimpleNamespace(task_id=task_id, start_from_trigger=True) + ), + ) + + dag_bag = mocker.Mock(spec=DBDagBag) + serialized_dags = { + stale_trigger.task_instance.dag_version_id: serialized_stale_dag, + healthy_trigger.task_instance.dag_version_id: serialized_healthy_dag, + } + dag_bag.get_serialized_dag_model.side_effect = lambda *, version_id, session: serialized_dags[version_id] + mocker.patch("airflow.jobs.triggerer_job_runner.DBDagBag", autospec=True, return_value=dag_bag) + mocker.patch( + "airflow.jobs.triggerer_job_runner.log_filename_template_renderer", + autospec=True, + return_value=lambda *, ti: "/logs/ti", + ) + mocker.patch( + "airflow.jobs.triggerer_job_runner.create_session", + autospec=True, + return_value=contextlib.nullcontext(mocker.sentinel.session), + ) + mocker.patch( + "airflow.jobs.triggerer_job_runner.TaskInstanceDTO.model_validate", + autospec=True, + side_effect=[mocker.Mock(spec=TaskInstanceDTO), mocker.Mock(spec=TaskInstanceDTO)], + ) + mocker.patch.object( + TriggerRunnerSupervisor, + "fetch_trigger_details", + autospec=True, + return_value={stale_trigger.id: stale_trigger, healthy_trigger.id: healthy_trigger}, + ) + mocker.patch.object( + TriggerRunnerSupervisor, + "fetch_non_task_trigger_ids", + autospec=True, + return_value=set(), + ) + + workloads = jobless_supervisor.build_trigger_workloads({stale_trigger.id, healthy_trigger.id}) + + workload_by_id = {workload.id: workload for workload in workloads} + assert workload_by_id.keys() == {stale_trigger.id, healthy_trigger.id} + assert workload_by_id[stale_trigger.id].dag_data is None + assert workload_by_id[healthy_trigger.id].dag_data == serialized_healthy_dag.data + assert workload_by_id[healthy_trigger.id].dag_run_data == {"dag_id": "example_dag"} + + def test_create_workload_sets_watched_assets_for_asset_only_trigger(jobless_supervisor, mocker): """_create_workload() should populate watched_assets when trigger.task_instance is None and assets exist.""" asset1 = mocker.Mock(spec=Asset) From f1e5967c3a358f98f32ca7e6cfc6af83ca176b4c Mon Sep 17 00:00:00 2001 From: Idan Izhaki Date: Mon, 13 Jul 2026 17:58:21 -0700 Subject: [PATCH 2/5] Document triggerer liveness fix Release notes should make the restored triggerer behavior visible to operators evaluating a 3.3 patch release. --- airflow-core/newsfragments/69842.bugfix.rst | 1 + 1 file changed, 1 insertion(+) create mode 100644 airflow-core/newsfragments/69842.bugfix.rst diff --git a/airflow-core/newsfragments/69842.bugfix.rst b/airflow-core/newsfragments/69842.bugfix.rst new file mode 100644 index 0000000000000..8d61335118517 --- /dev/null +++ b/airflow-core/newsfragments/69842.bugfix.rst @@ -0,0 +1 @@ +Keep the triggerer running when a deferred task is missing from its serialized Dag. From 16e34164d7505f31f0782f61fc6be7abdfc54d23 Mon Sep 17 00:00:00 2001 From: Idan Izhaki Date: Mon, 13 Jul 2026 18:14:30 -0700 Subject: [PATCH 3/5] Persist trigger context requirements --- airflow-core/docs/migrations-ref.rst | 4 +- .../execution_api/routes/task_instances.py | 1 + .../src/airflow/dag_processing/collection.py | 1 + .../src/airflow/jobs/triggerer_job_runner.py | 81 ++++--- ...3_4_0_add_start_from_trigger_to_trigger.py | 51 +++++ .../src/airflow/models/taskinstance.py | 1 + airflow-core/src/airflow/models/trigger.py | 7 +- airflow-core/src/airflow/utils/db.py | 2 +- .../versions/head/test_task_instances.py | 1 + .../unit/dag_processing/test_collection.py | 1 + .../tests/unit/jobs/test_triggerer_job.py | 210 +++++++++++++----- .../tests/unit/models/test_callback.py | 1 + .../tests/unit/models/test_taskinstance.py | 1 + .../tests/unit/models/test_trigger.py | 22 ++ 14 files changed, 294 insertions(+), 90 deletions(-) create mode 100644 airflow-core/src/airflow/migrations/versions/0127_3_4_0_add_start_from_trigger_to_trigger.py diff --git a/airflow-core/docs/migrations-ref.rst b/airflow-core/docs/migrations-ref.rst index 3dd6e96a7d00b..8a44e9b5c7080 100644 --- a/airflow-core/docs/migrations-ref.rst +++ b/airflow-core/docs/migrations-ref.rst @@ -39,7 +39,9 @@ Here's the list of all the Database Migrations that are executed via when you ru +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | Revision ID | Revises ID | Airflow Version | Description | +=========================+==================+===================+==============================================================+ -| ``c4e7a1f9b2d0`` (head) | ``436dc127462c`` | ``3.4.0`` | Add index on asset.uri. | +| ``9f3c2d4a1b7e`` (head) | ``c4e7a1f9b2d0`` | ``3.4.0`` | Add start_from_trigger to trigger. | ++-------------------------+------------------+-------------------+--------------------------------------------------------------+ +| ``c4e7a1f9b2d0`` | ``436dc127462c`` | ``3.4.0`` | Add index on asset.uri. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | ``436dc127462c`` | ``5a5d3253e946`` | ``3.4.0`` | Drop span_status column. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ diff --git a/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py b/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py index c1bac7960234d..8447aba4922b1 100644 --- a/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py +++ b/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py @@ -688,6 +688,7 @@ def _create_ti_state_update_query_and_update_state( kwargs={}, queue=ti_patch_payload.queue, team_name=get_team_name_for_ti(task_instance_id, session), + start_from_trigger=False, ) trigger_row.encrypted_kwargs = trigger_kwargs session.add(trigger_row) diff --git a/airflow-core/src/airflow/dag_processing/collection.py b/airflow-core/src/airflow/dag_processing/collection.py index 8fa5209ff15fb..3897cfa3a769f 100644 --- a/airflow-core/src/airflow/dag_processing/collection.py +++ b/airflow-core/src/airflow/dag_processing/collection.py @@ -1186,6 +1186,7 @@ def add_asset_trigger_references( classpath=triggers[trigger_hash]["classpath"], kwargs=triggers[trigger_hash]["kwargs"], team_name=team_name, + start_from_trigger=False, ) for trigger_hash in all_trigger_hashes if trigger_hash not in orm_triggers diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 27318adab63ad..0e71de22bfc04 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -134,6 +134,7 @@ from airflow.utils.helpers import log_filename_template_renderer, prune_dict from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.session import create_session, provide_session +from airflow.utils.state import TaskInstanceState if TYPE_CHECKING: from opentelemetry.util._decorator import _AgnosticContextManager @@ -874,50 +875,64 @@ def _create_workload( ti=ser_ti, # type: ignore ) + workload = workloads.RunTrigger( + id=trigger.id, + classpath=trigger.classpath, + encrypted_kwargs=trigger.encrypted_kwargs, + ti=ser_ti, + timeout_after=trigger.task_instance.trigger_timeout, + ) + if trigger.start_from_trigger is False: + return workload + serialized_dag_model = dag_bag.get_serialized_dag_model( version_id=trigger.task_instance.dag_version_id, session=session, ) - if serialized_dag_model: - try: - task = serialized_dag_model.dag.get_task(trigger.task_instance.task_id) - except TaskNotFound: - # Avoid this lookup for ordinary deferrals once trigger context requirements are persisted; - # tracked at https://github.com/apache/airflow/issues/69841 - log.warning( - "Task for deferred trigger was not found in serialized Dag; " - "starting trigger without Dag context", - trigger_id=trigger.id, - ti_id=trigger.task_instance.id, - dag_id=trigger.task_instance.dag_id, - task_id=trigger.task_instance.task_id, - dag_version_id=trigger.task_instance.dag_version_id, - ) - task = None - - # When a TaskInstance of a Trigger contains a task with start_from_trigger enabled, - # it means we need to load the SerializedDagModel so we can build a RuntimeTaskInstance later on which - # will allow us to build a context on which we will render the templated fields. - if task is not None and task.start_from_trigger: - log.info("Start from trigger enabled for task %s", task.task_id) - dag_run = trigger.task_instance.get_dagrun(session=session) - - return workloads.RunTrigger( - id=trigger.id, - classpath=trigger.classpath, - encrypted_kwargs=trigger.encrypted_kwargs, - ti=ser_ti, - timeout_after=trigger.task_instance.trigger_timeout, - dag_data=serialized_dag_model.data, - dag_run_data=dag_run.dag_run_data.model_dump(exclude_unset=True), - ) + if serialized_dag_model is None: + log.warning( + "Serialized Dag for context-required trigger was not found; skipping trigger", + trigger_id=trigger.id, + ti_id=trigger.task_instance.id, + dag_id=trigger.task_instance.dag_id, + task_id=trigger.task_instance.task_id, + dag_version_id=trigger.task_instance.dag_version_id, + ) + return None + + try: + task = serialized_dag_model.dag.get_task(trigger.task_instance.task_id) + except TaskNotFound: + trigger.task_instance.state = TaskInstanceState.REMOVED + trigger.task_instance.trigger_id = None + log.warning( + "Task for deferred trigger was not found in serialized Dag; removing TaskInstance", + trigger_id=trigger.id, + ti_id=trigger.task_instance.id, + dag_id=trigger.task_instance.dag_id, + task_id=trigger.task_instance.task_id, + dag_version_id=trigger.task_instance.dag_version_id, + ) + return None + + if trigger.start_from_trigger is None: + trigger.start_from_trigger = task.start_from_trigger + + if trigger.start_from_trigger is False: + return workload + + log.info("Start from trigger enabled for task %s", task.task_id) + dag_run = trigger.task_instance.get_dagrun(session=session) + return workloads.RunTrigger( id=trigger.id, classpath=trigger.classpath, encrypted_kwargs=trigger.encrypted_kwargs, ti=ser_ti, timeout_after=trigger.task_instance.trigger_timeout, + dag_data=serialized_dag_model.data, + dag_run_data=dag_run.dag_run_data.model_dump(exclude_unset=True), ) def fetch_trigger_details(self, trigger_ids: set[int], *, session: Session) -> dict[int, Trigger]: diff --git a/airflow-core/src/airflow/migrations/versions/0127_3_4_0_add_start_from_trigger_to_trigger.py b/airflow-core/src/airflow/migrations/versions/0127_3_4_0_add_start_from_trigger_to_trigger.py new file mode 100644 index 0000000000000..c29075ab1eca0 --- /dev/null +++ b/airflow-core/src/airflow/migrations/versions/0127_3_4_0_add_start_from_trigger_to_trigger.py @@ -0,0 +1,51 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +""" +Add start_from_trigger to trigger. + +Revision ID: 9f3c2d4a1b7e +Revises: c4e7a1f9b2d0 +Create Date: 2026-07-13 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +from airflow.migrations.utils import disable_sqlite_fkeys + +# revision identifiers, used by Alembic. +revision = "9f3c2d4a1b7e" +down_revision = "c4e7a1f9b2d0" +branch_labels = None +depends_on = None +airflow_version = "3.4.0" + + +def upgrade(): + """Add start_from_trigger to trigger.""" + with disable_sqlite_fkeys(op): + with op.batch_alter_table("trigger", schema=None) as batch_op: + batch_op.add_column(sa.Column("start_from_trigger", sa.Boolean(), nullable=True)) + + +def downgrade(): + """Remove start_from_trigger from trigger.""" + with disable_sqlite_fkeys(op): + with op.batch_alter_table("trigger", schema=None) as batch_op: + batch_op.drop_column("start_from_trigger") diff --git a/airflow-core/src/airflow/models/taskinstance.py b/airflow-core/src/airflow/models/taskinstance.py index 9f8eb0805c4f4..6515314ed448c 100644 --- a/airflow-core/src/airflow/models/taskinstance.py +++ b/airflow-core/src/airflow/models/taskinstance.py @@ -1784,6 +1784,7 @@ def defer_task(self, *, session: Session = NEW_SESSION) -> bool: classpath=start_trigger_args.trigger_cls, kwargs=trigger_kwargs, team_name=team_name, + start_from_trigger=True, ) # First, make the trigger entry diff --git a/airflow-core/src/airflow/models/trigger.py b/airflow-core/src/airflow/models/trigger.py index 386a6c931fcc6..fc01f15ca4a3e 100644 --- a/airflow-core/src/airflow/models/trigger.py +++ b/airflow-core/src/airflow/models/trigger.py @@ -24,7 +24,7 @@ from traceback import format_exception from typing import TYPE_CHECKING, Any -from sqlalchemy import ForeignKey, Integer, String, Text, delete, func, or_, select, update +from sqlalchemy import Boolean, ForeignKey, Integer, String, Text, delete, func, or_, select, update from sqlalchemy.ext.associationproxy import association_proxy from sqlalchemy.orm import Mapped, Session, mapped_column, relationship, selectinload from sqlalchemy.sql.functions import coalesce @@ -99,6 +99,7 @@ class Trigger(Base): created_date: Mapped[datetime.datetime] = mapped_column(UtcDateTime, nullable=False) triggerer_id: Mapped[int | None] = mapped_column(Integer, nullable=True) queue: Mapped[str | None] = mapped_column(String(256), nullable=True) + start_from_trigger: Mapped[bool | None] = mapped_column(Boolean, nullable=True) # Denormalized from dag_bundle_team to keep the triggerer's ~1s polling queries join-free, # especially since it's eventually consistent and trigger rows are ephemeral. @@ -132,6 +133,7 @@ def __init__( created_date: datetime.datetime | None = None, queue: str | None = None, team_name: str | None = None, + start_from_trigger: bool | None = False, ) -> None: super().__init__() self.classpath = classpath @@ -139,6 +141,7 @@ def __init__( self.created_date = created_date or timezone.utcnow() self.queue = queue self.team_name = team_name + self.start_from_trigger = start_from_trigger @property def kwargs(self) -> dict[str, Any]: @@ -200,7 +203,7 @@ def rotate_fernet_key(self): def from_object(cls, trigger: BaseTrigger) -> Trigger: """Alternative constructor that creates a trigger row based directly off of a Trigger object.""" classpath, kwargs = trigger.serialize() - return cls(classpath=classpath, kwargs=kwargs) + return cls(classpath=classpath, kwargs=kwargs, start_from_trigger=False) @classmethod @provide_session diff --git a/airflow-core/src/airflow/utils/db.py b/airflow-core/src/airflow/utils/db.py index ce4ef036d20fd..a478634db5ea3 100644 --- a/airflow-core/src/airflow/utils/db.py +++ b/airflow-core/src/airflow/utils/db.py @@ -117,7 +117,7 @@ class MappedClassProtocol(Protocol): "3.1.8": "509b94a1042d", "3.2.0": "1d6611b6ab7c", "3.3.0": "d2f4e1b3c5a7", - "3.4.0": "c4e7a1f9b2d0", + "3.4.0": "9f3c2d4a1b7e", } # Prefix used to identify tables holding data moved during migration. diff --git a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py index 542ce7eaaf15b..ac1e54302563b 100644 --- a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py +++ b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py @@ -1597,6 +1597,7 @@ def test_ti_update_state_to_deferred( "key": "value", "moment": datetime(2024, 12, 18, 00, 00, 1, tzinfo=timezone.utc), } + assert t[0].start_from_trigger is False if queues_enabled: assert t[0].queue == "default" else: diff --git a/airflow-core/tests/unit/dag_processing/test_collection.py b/airflow-core/tests/unit/dag_processing/test_collection.py index ad21350be6bcb..1a046811a0f5c 100644 --- a/airflow-core/tests/unit/dag_processing/test_collection.py +++ b/airflow-core/tests/unit/dag_processing/test_collection.py @@ -394,6 +394,7 @@ def test_add_asset_trigger_references_populates_team_name( triggers = session.scalars(select(Trigger)).all() assert len(triggers) == 1 assert triggers[0].team_name == expected + assert triggers[0].start_from_trigger is False @pytest.mark.usefixtures("testing_dag_bundle") def test_add_asset_trigger_references_hash_consistency(self, dag_maker, session): diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 3866c563e42a0..0a2e8f78cae53 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -101,6 +101,7 @@ _ResponseFrame, ) from airflow.sdk.execution_time.context import AssetStateStoreAccessors +from airflow.serialization.definitions.dag import SerializedDAG from airflow.serialization.serialized_objects import LazyDeserializedDAG from airflow.triggers.base import BaseEventTrigger, BaseTrigger, TriggerEvent from airflow.triggers.shared_stream import SharedStreamProducer @@ -612,21 +613,64 @@ def test_create_workload_uses_supervisor_id_without_job(jobless_supervisor, mock assert factory.log_path == f"/logs/ti.trigger.{jobless_supervisor.id}.log" -def test_create_workload_missing_task_returns_context_free_workload(jobless_supervisor, mocker, caplog): +def _make_task_trigger(*, trigger_id, start_from_trigger, task_id="task", dag_version_id=None): ti = SimpleNamespace( id=uuid.uuid4(), dag_id="example_dag", - dag_version_id=uuid.uuid4(), - task_id="removed_task", + dag_version_id=dag_version_id or uuid.uuid4(), + task_id=task_id, trigger_timeout=None, + state=TaskInstanceState.DEFERRED, + trigger_id=trigger_id, + get_dagrun=lambda *, session: SimpleNamespace( + dag_run_data=SimpleNamespace( + model_dump=lambda **kwargs: {"dag_id": "example_dag"}, + ) + ), ) - trigger = SimpleNamespace( - id=7, + return SimpleNamespace( + id=trigger_id, classpath="some.path.Trigger", encrypted_kwargs="encrypted", task_instance=ti, + start_from_trigger=start_from_trigger, ) + +def _make_serialized_dag(mocker, *, task_id="task", start_from_trigger=False, missing_task=False): + dag = mocker.Mock(spec=SerializedDAG) + if missing_task: + dag.get_task.side_effect = TaskNotFound(task_id) + else: + dag.get_task.return_value = SimpleNamespace( + task_id=task_id, + start_from_trigger=start_from_trigger, + ) + serialized_dag = mocker.Mock(spec=SerializedDagModel) + serialized_dag.dag = dag + serialized_dag.data = {"dag": "serialized"} + return serialized_dag + + +@pytest.mark.parametrize( + ("stored_requirement", "task_requirement", "expect_lookup", "expected_requirement"), + [ + pytest.param(False, False, False, False, id="ordinary"), + pytest.param(True, True, True, True, id="direct-to-triggerer"), + pytest.param(None, False, True, False, id="legacy-ordinary"), + pytest.param(None, True, True, True, id="legacy-direct-to-triggerer"), + ], +) +def test_create_workload_uses_persisted_context_requirement( + jobless_supervisor, + mocker, + stored_requirement, + task_requirement, + expect_lookup, + expected_requirement, +): + trigger = _make_task_trigger(trigger_id=7, start_from_trigger=stored_requirement) + ser_ti = mocker.Mock(spec=TaskInstanceDTO) mocker.patch( "airflow.jobs.triggerer_job_runner.TaskInstanceDTO.model_validate", @@ -634,10 +678,60 @@ def test_create_workload_missing_task_returns_context_free_workload(jobless_supe return_value=ser_ti, ) - def get_task(task_id): - raise TaskNotFound(task_id) + serialized_dag_model = _make_serialized_dag( + mocker, + task_id=trigger.task_instance.task_id, + start_from_trigger=task_requirement, + ) + dag_bag = mocker.Mock(spec=DBDagBag) + dag_bag.get_serialized_dag_model.return_value = serialized_dag_model + + workload = jobless_supervisor._create_workload( + trigger=trigger, + dag_bag=dag_bag, + render_log_fname=lambda *, ti: "/logs/ti", + session=mocker.sentinel.session, + ) + + assert workload is not None + assert workload.id == trigger.id + assert workload.ti == ser_ti + assert trigger.start_from_trigger is expected_requirement + if expect_lookup: + dag_bag.get_serialized_dag_model.assert_called_once_with( + version_id=trigger.task_instance.dag_version_id, + session=mocker.sentinel.session, + ) + serialized_dag_model.dag.get_task.assert_called_once_with(trigger.task_instance.task_id) + else: + dag_bag.get_serialized_dag_model.assert_not_called() + + if expected_requirement: + assert workload.dag_data == serialized_dag_model.data + assert workload.dag_run_data == {"dag_id": "example_dag"} + else: + assert workload.dag_data is None + assert workload.dag_run_data is None + + +@pytest.mark.parametrize("stored_requirement", [True, None]) +def test_create_workload_missing_task_removes_ti(jobless_supervisor, mocker, caplog, stored_requirement): + trigger = _make_task_trigger( + trigger_id=7, + start_from_trigger=stored_requirement, + task_id="removed_task", + ) + mocker.patch( + "airflow.jobs.triggerer_job_runner.TaskInstanceDTO.model_validate", + autospec=True, + return_value=mocker.Mock(spec=TaskInstanceDTO), + ) - serialized_dag_model = SimpleNamespace(dag=SimpleNamespace(get_task=get_task)) + serialized_dag_model = _make_serialized_dag( + mocker, + task_id=trigger.task_instance.task_id, + missing_task=True, + ) dag_bag = mocker.Mock(spec=DBDagBag) dag_bag.get_serialized_dag_model.return_value = serialized_dag_model @@ -648,15 +742,45 @@ def get_task(task_id): session=mocker.sentinel.session, ) - assert workload == RunTrigger( - id=trigger.id, - classpath=trigger.classpath, - encrypted_kwargs=trigger.encrypted_kwargs, - ti=ser_ti, + assert workload is None + assert trigger.task_instance.state == TaskInstanceState.REMOVED + assert trigger.task_instance.trigger_id is None + assert { + "event": "Task for deferred trigger was not found in serialized Dag; removing TaskInstance", + "trigger_id": trigger.id, + "ti_id": trigger.task_instance.id, + "dag_id": trigger.task_instance.dag_id, + "task_id": trigger.task_instance.task_id, + "dag_version_id": trigger.task_instance.dag_version_id, + } in caplog + + +@pytest.mark.parametrize("stored_requirement", [True, None]) +def test_create_workload_missing_serialized_dag_retries( + jobless_supervisor, mocker, caplog, stored_requirement +): + trigger = _make_task_trigger(trigger_id=7, start_from_trigger=stored_requirement) + mocker.patch( + "airflow.jobs.triggerer_job_runner.TaskInstanceDTO.model_validate", + autospec=True, + return_value=mocker.Mock(spec=TaskInstanceDTO), ) + + dag_bag = mocker.Mock(spec=DBDagBag) + dag_bag.get_serialized_dag_model.return_value = None + + workload = jobless_supervisor._create_workload( + trigger=trigger, + dag_bag=dag_bag, + render_log_fname=lambda *, ti: "/logs/ti", + session=mocker.sentinel.session, + ) + + assert workload is None + assert trigger.task_instance.state == TaskInstanceState.DEFERRED + assert trigger.task_instance.trigger_id == trigger.id assert { - "event": "Task for deferred trigger was not found in serialized Dag; " - "starting trigger without Dag context", + "event": "Serialized Dag for context-required trigger was not found; skipping trigger", "trigger_id": trigger.id, "ti_id": trigger.task_instance.id, "dag_id": trigger.task_instance.dag_id, @@ -666,47 +790,26 @@ def get_task(task_id): def test_build_trigger_workloads_continues_after_missing_task(jobless_supervisor, mocker): - def get_missing_task(task_id): - raise TaskNotFound(task_id) - - stale_ti = SimpleNamespace( - id=uuid.uuid4(), - dag_id="example_dag", - dag_version_id=uuid.uuid4(), + stale_trigger = _make_task_trigger( + trigger_id=1, + start_from_trigger=True, task_id="removed_task", - trigger_timeout=None, ) - stale_trigger = SimpleNamespace( - id=1, - classpath="some.path.StaleTrigger", - encrypted_kwargs="stale", - task_instance=stale_ti, - ) - healthy_ti = SimpleNamespace( - id=uuid.uuid4(), - dag_id="example_dag", - dag_version_id=uuid.uuid4(), + healthy_trigger = _make_task_trigger( + trigger_id=2, + start_from_trigger=True, task_id="healthy_task", - trigger_timeout=None, - get_dagrun=lambda *, session: SimpleNamespace( - dag_run_data=SimpleNamespace( - model_dump=lambda **kwargs: {"dag_id": "example_dag"}, - ) - ), - ) - healthy_trigger = SimpleNamespace( - id=2, - classpath="some.path.HealthyTrigger", - encrypted_kwargs="healthy", - task_instance=healthy_ti, ) - serialized_stale_dag = SimpleNamespace(dag=SimpleNamespace(get_task=get_missing_task)) - serialized_healthy_dag = SimpleNamespace( - data={"dag": "serialized"}, - dag=SimpleNamespace( - get_task=lambda task_id: SimpleNamespace(task_id=task_id, start_from_trigger=True) - ), + serialized_stale_dag = _make_serialized_dag( + mocker, + task_id=stale_trigger.task_instance.task_id, + missing_task=True, + ) + serialized_healthy_dag = _make_serialized_dag( + mocker, + task_id=healthy_trigger.task_instance.task_id, + start_from_trigger=True, ) dag_bag = mocker.Mock(spec=DBDagBag) @@ -747,10 +850,11 @@ def get_missing_task(task_id): workloads = jobless_supervisor.build_trigger_workloads({stale_trigger.id, healthy_trigger.id}) workload_by_id = {workload.id: workload for workload in workloads} - assert workload_by_id.keys() == {stale_trigger.id, healthy_trigger.id} - assert workload_by_id[stale_trigger.id].dag_data is None + assert workload_by_id.keys() == {healthy_trigger.id} assert workload_by_id[healthy_trigger.id].dag_data == serialized_healthy_dag.data assert workload_by_id[healthy_trigger.id].dag_run_data == {"dag_id": "example_dag"} + assert stale_trigger.task_instance.state == TaskInstanceState.REMOVED + assert stale_trigger.task_instance.trigger_id is None def test_create_workload_sets_watched_assets_for_asset_only_trigger(jobless_supervisor, mocker): diff --git a/airflow-core/tests/unit/models/test_callback.py b/airflow-core/tests/unit/models/test_callback.py index b2b13582710fc..73b6341a94907 100644 --- a/airflow-core/tests/unit/models/test_callback.py +++ b/airflow-core/tests/unit/models/test_callback.py @@ -185,6 +185,7 @@ def test_queue(self, session): assert isinstance(callback.trigger, Trigger) assert callback.trigger.kwargs["callback_path"] == TEST_ASYNC_CALLBACK.path assert callback.trigger.kwargs["callback_kwargs"] == TEST_ASYNC_CALLBACK.kwargs + assert callback.trigger.start_from_trigger is False assert callback.state == CallbackState.QUEUED @staticmethod diff --git a/airflow-core/tests/unit/models/test_taskinstance.py b/airflow-core/tests/unit/models/test_taskinstance.py index f2f5675ae566c..fb7b02d8e4a88 100644 --- a/airflow-core/tests/unit/models/test_taskinstance.py +++ b/airflow-core/tests/unit/models/test_taskinstance.py @@ -2922,6 +2922,7 @@ def test_defer_task(create_task_instance): assert isinstance(trigger_row, Trigger) assert trigger_row.classpath == "trigger_cls" assert trigger_row.kwargs == {"key": "value"} + assert trigger_row.start_from_trigger is True # Check that session.flush was called session.flush.assert_called_once() diff --git a/airflow-core/tests/unit/models/test_trigger.py b/airflow-core/tests/unit/models/test_trigger.py index 91a6a92e27cb3..9098b97a33b5a 100644 --- a/airflow-core/tests/unit/models/test_trigger.py +++ b/airflow-core/tests/unit/models/test_trigger.py @@ -97,6 +97,28 @@ def test_trigger_team_name_stored(session, testing_team): assert loaded.team_name == "testing" +@pytest.mark.parametrize( + ("constructor_options", "expected"), + [ + pytest.param({}, False, id="default"), + pytest.param({"start_from_trigger": None}, None, id="legacy-null"), + pytest.param({"start_from_trigger": True}, True, id="direct-to-triggerer"), + ], +) +def test_trigger_start_from_trigger_stored(session, constructor_options, expected): + trigger = Trigger( + classpath="airflow.triggers.testing.SuccessTrigger", + kwargs={}, + **constructor_options, + ) + session.add(trigger) + session.flush() + trigger_id = trigger.id + session.expire_all() + + assert session.get(Trigger, trigger_id).start_from_trigger is expected + + def test_fetch_trigger_ids_with_non_task_associations(session): # Create triggers asset_trigger = Trigger(classpath="airflow.triggers.testing.SuccessTrigger1", kwargs={}) From f1462fe0cc42df8197e3af53e56cebb29040aedd Mon Sep 17 00:00:00 2001 From: Idan Izhaki Date: Mon, 13 Jul 2026 18:15:43 -0700 Subject: [PATCH 4/5] Document trigger context persistence --- airflow-core/newsfragments/69843.improvement.rst | 1 + 1 file changed, 1 insertion(+) create mode 100644 airflow-core/newsfragments/69843.improvement.rst diff --git a/airflow-core/newsfragments/69843.improvement.rst b/airflow-core/newsfragments/69843.improvement.rst new file mode 100644 index 0000000000000..f42a20ccc7344 --- /dev/null +++ b/airflow-core/newsfragments/69843.improvement.rst @@ -0,0 +1 @@ +Avoid loading serialized Dags for ordinary deferred task triggers by persisting whether Dag context is required. From 4da7203766ddd4172a5cd8a0fc1337dbb03379b9 Mon Sep 17 00:00:00 2001 From: Idan Izhaki Date: Mon, 13 Jul 2026 18:23:34 -0700 Subject: [PATCH 5/5] Test legacy trigger context persistence --- .../tests/unit/jobs/test_triggerer_job.py | 44 ++++++++++++++++++- 1 file changed, 43 insertions(+), 1 deletion(-) diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 0a2e8f78cae53..35b3f53a8fb52 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -47,6 +47,7 @@ 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 update from structlog.typing import FilteringBoundLogger from airflow._shared.timezones import timezone @@ -103,7 +104,7 @@ from airflow.sdk.execution_time.context import AssetStateStoreAccessors from airflow.serialization.definitions.dag import SerializedDAG from airflow.serialization.serialized_objects import LazyDeserializedDAG -from airflow.triggers.base import BaseEventTrigger, BaseTrigger, TriggerEvent +from airflow.triggers.base import BaseEventTrigger, BaseTrigger, StartTriggerArgs, TriggerEvent from airflow.triggers.shared_stream import SharedStreamProducer from airflow.triggers.testing import FailureTrigger, SuccessTrigger from airflow.utils.state import State, TaskInstanceState @@ -714,6 +715,47 @@ def test_create_workload_uses_persisted_context_requirement( assert workload.dag_run_data is None +@pytest.mark.parametrize("task_requirement", [False, True]) +def test_build_trigger_workloads_persists_legacy_context_requirement( + session, + supervisor_builder, + task_requirement, +): + class ContextRequiredOperator(BaseOperator): + start_from_trigger = True + start_trigger_args = StartTriggerArgs( + trigger_cls="airflow.triggers.testing.SuccessTrigger", + trigger_kwargs={}, + next_method="execute_complete", + ) + + operator_class = ContextRequiredOperator if task_requirement else BaseOperator + operator = operator_class(task_id="test_ti") + _, _, trigger, task_instance = create_trigger_in_db( + session, + TimeDeltaTrigger(datetime.timedelta(days=7)), + operator=operator, + ) + serialized_dag = DBDagBag().get_serialized_dag_model( + version_id=task_instance.dag_version_id, + session=session, + ) + assert serialized_dag.dag.get_task(task_instance.task_id).start_from_trigger is task_requirement + trigger_id = trigger.id + session.execute( + update(Trigger).where(Trigger.id == trigger_id).values(start_from_trigger=None), + ) + session.commit() + session.expire_all() + assert session.get(Trigger, trigger_id).start_from_trigger is None + + built_workloads = supervisor_builder().build_trigger_workloads({trigger_id}) + session.expire_all() + + assert len(built_workloads) == 1 + assert session.get(Trigger, trigger_id).start_from_trigger is task_requirement + + @pytest.mark.parametrize("stored_requirement", [True, None]) def test_create_workload_missing_task_removes_ti(jobless_supervisor, mocker, caplog, stored_requirement): trigger = _make_task_trigger(