Skip to content
Merged
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/69821.bugfix.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Deferrable tasks that fail via a trigger-emitted ``TaskFailedEvent`` now respect retries: if the task has retries remaining it goes ``up_for_retry`` and runs ``on_retry_callback``, instead of always failing terminally and running ``on_failure_callback``.
46 changes: 40 additions & 6 deletions airflow-core/src/airflow/models/trigger.py
Original file line number Diff line number Diff line change
Expand Up @@ -561,13 +561,35 @@ def _(event: BaseTaskEndEvent, *, task_instance: TaskInstance, session: Session)
from airflow.callbacks.database_callback_sink import DatabaseCallbackSink
from airflow.utils.state import TaskInstanceState

# Mark the task with terminal state and prevent it from resuming on worker
# Prevent the task from resuming on a worker.
task_instance.trigger_id = None
task_instance.set_state(event.task_instance_state, session=session)

callback_type = event.task_instance_state
should_retry = False

if event.task_instance_state == TaskInstanceState.FAILED:
# Load the serialized task so retry eligibility matches the normal task path.
try:
from airflow.models.dagbag import DBDagBag

dag = DBDagBag().get_dag_for_run(dag_run=task_instance.dag_run, session=session)
if dag is not None:
task_instance.task = dag.get_task(task_instance.task_id)
should_retry = task_instance.is_eligible_to_retry()
except Exception:
log.exception(
"Could not load task for %s; failing terminally without retry routing", task_instance
)
if should_retry:
callback_type = TaskInstanceState.UP_FOR_RETRY

def _submit_callback_if_necessary() -> None:
"""Submit a callback request if the task state is SUCCESS or FAILED."""
if event.task_instance_state in (TaskInstanceState.SUCCESS, TaskInstanceState.FAILED):
"""Submit a callback request if the task state is SUCCESS, FAILED, or UP_FOR_RETRY."""
if callback_type in (
TaskInstanceState.SUCCESS,
TaskInstanceState.FAILED,
TaskInstanceState.UP_FOR_RETRY,
):
if task_instance.dag_model.relative_fileloc is None:
raise RuntimeError("relative_fileloc should not be None for a finished task")
from airflow.models.dag_version import _resolve_version_data
Expand All @@ -591,7 +613,7 @@ def _submit_callback_if_necessary() -> None:
request = TaskCallbackRequest(
filepath=task_instance.dag_model.relative_fileloc,
ti=task_instance,
task_callback_type=event.task_instance_state,
task_callback_type=callback_type,
bundle_name=bundle_name,
bundle_version=bundle_version,
version_data=version_data,
Expand All @@ -604,10 +626,22 @@ def _submit_callback_if_necessary() -> None:

def _push_xcoms_if_necessary() -> None:
"""Pushes XComs to the database if they are provided."""
if event.xcoms:
if event.xcoms and callback_type != TaskInstanceState.UP_FOR_RETRY:
for key, value in event.xcoms.items():
task_instance.xcom_push(key=key, value=value)

# Send the callback before mutating task state so it reflects the retry-vs-terminal
# decision derived above.
_submit_callback_if_necessary()

if should_retry:
task_instance.end_date = timezone.utcnow()
task_instance.set_duration()
task_instance.clear_next_method_args()
task_instance.prepare_db_for_next_try(session)
task_instance.state = TaskInstanceState.UP_FOR_RETRY
else:
task_instance.set_state(event.task_instance_state, session=session)

_push_xcoms_if_necessary()
session.flush()
76 changes: 74 additions & 2 deletions airflow-core/tests/unit/models/test_trigger.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
from airflow.models import TaskInstance, Trigger
from airflow.models.asset import AssetEvent, AssetModel, AssetWatcherModel
from airflow.models.callback import Callback, TriggererCallback
from airflow.models.taskinstancehistory import TaskInstanceHistory
from airflow.models.xcom import XComModel
from airflow.providers.standard.operators.empty import EmptyOperator
from airflow.sdk.definitions.callback import AsyncCallback
Expand All @@ -48,7 +49,7 @@
TriggerEvent,
)
from airflow.utils.session import create_session
from airflow.utils.state import State
from airflow.utils.state import State, TaskInstanceState

from tests_common.test_utils.asserts import assert_queries_count
from tests_common.test_utils.config import conf_vars
Expand Down Expand Up @@ -296,11 +297,15 @@ def test_submit_event_task_end(mock_utcnow, session, create_task_instance, event
# Make a trigger
trigger = Trigger(classpath="does.not.matter", kwargs={})
session.add(trigger)
# Make a TaskInstance that's deferred and waiting on it
# Make a TaskInstance that's deferred and waiting on it. A deferred task has
# already started running, so it has a start_date; set one so duration can be
# computed. Unlike set_state, handle_failure (used by the FAILED path) does not
# synthesize a missing start_date, matching the scheduler executor-event path.
task_instance = create_task_instance(
session=session, logical_date=timezone.utcnow(), state=State.DEFERRED
)
task_instance.trigger_id = trigger.id
task_instance.start_date = now.subtract(seconds=10)
session.commit()

def get_xcoms(ti):
Expand Down Expand Up @@ -362,6 +367,73 @@ def test_submit_event_task_end_callback_includes_version_data(mock_send, session
assert request.version_data == version_data


@pytest.mark.parametrize(
("retries", "expected_state", "expected_callback_type", "expect_history_row"),
[
(1, TaskInstanceState.UP_FOR_RETRY, TaskInstanceState.UP_FOR_RETRY, True),
(0, TaskInstanceState.FAILED, TaskInstanceState.FAILED, False),
],
)
@patch("airflow.callbacks.database_callback_sink.DatabaseCallbackSink.send")
def test_submit_event_task_end_failed_respects_retries(
mock_send,
session,
create_task_instance,
retries,
expected_state,
expected_callback_type,
expect_history_row,
):
"""A trigger-emitted TaskFailedEvent should respect retry-eligibility: a deferred task with
retries remaining goes UP_FOR_RETRY (on_retry_callback), not straight to FAILED.

On the retry path, the finished try must also be archived to task_instance_history so
prior-try log lookups keep working after the trigger ends the deferred try.
"""
trigger = Trigger(classpath="does.not.matter", kwargs={})
session.add(trigger)
task_instance = create_task_instance(
session=session,
logical_date=timezone.utcnow(),
state=State.DEFERRED,
default_args={"retries": retries},
)
task_instance.trigger_id = trigger.id
task_instance.try_number = 1
task_instance.max_tries = retries
old_ti_id = task_instance.id
session.commit()

Trigger.submit_event(trigger.id, TaskFailedEvent(), session=session)
session.flush()

ti = session.scalar(select(TaskInstance))
assert ti.state == expected_state

mock_send.assert_called_once()
request = mock_send.call_args.kwargs["callback"]
assert request.task_callback_type == expected_callback_type

assert ti.next_method is None
assert ti.next_kwargs is None
assert ti.end_date is not None

tih = session.scalars(
select(TaskInstanceHistory).where(
TaskInstanceHistory.dag_id == ti.dag_id,
TaskInstanceHistory.task_id == ti.task_id,
TaskInstanceHistory.run_id == ti.run_id,
)
).all()
if expect_history_row:
assert len(tih) == 1
assert ti.id != old_ti_id
assert tih[0].task_instance_id == old_ti_id
else:
assert tih == []
assert ti.id == old_ti_id


@pytest.fixture
def create_triggerer():
"""Fixture factory which creates individual test Triggerer instances."""
Expand Down