Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions airflow-core/newsfragments/69842.bugfix.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Keep the triggerer running when a deferred task is missing from its serialized Dag.
19 changes: 17 additions & 2 deletions airflow-core/src/airflow/jobs/triggerer_job_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
144 changes: 144 additions & 0 deletions airflow-core/tests/unit/jobs/test_triggerer_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down