Skip to content

Commit 08d8d0c

Browse files
committed
Run loop gates and advance iterations atomically
The scheduler, the API and the UI must never see every task of a loop finished while its next pass does not exist yet. If we recorded a continuing gate as successful and created the next pass afterwards, a crash between the two would leave recovery unable to tell an unfinished advancement from a completed decision. Gate completion and next-pass creation therefore happen in one transaction under the DagRun lock, which also stops a duplicate completion from advancing the loop twice. The continue-or-stop decision is a reserved XCom stored under the gate's attempt UUID. A retried gate starts without a decision, and a retired attempt cannot advance the loop. An invalid decision moves the gate attempt to retry or failed within the same request; rejecting the request instead would leave the attempt running with nobody to finish it. Gates use normal task dependencies and an explicit loop context. A barrier that waited for every body task would override trigger rules that allow early progress. Workers and dag.test supply the same coordinates, so decisions and previous-iteration reads reach the intended gate and data. Skipping downstream of a task from outside a loop skips every live pass. Skipping only one would let later passes run work the author meant to skip.
1 parent c3b6ed7 commit 08d8d0c

50 files changed

Lines changed: 3094 additions & 186 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎airflow-core/docs/authoring-and-scheduling/loops.rst‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -62,7 +62,7 @@ iterations.
6262
This example improves an estimate of the square root of two until the error
6363
is small enough:
6464

65-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
65+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
6666
:start-after: [START refine_estimate]
6767
:end-before: [END refine_estimate]
6868

@@ -92,15 +92,15 @@ For a fixed-count loop, omit ``until``. The definition ``refine.loop(max_iterati
9292
three iterations, carrying results between them. Reaching the cap completes
9393
a fixed-count loop successfully. Its gate is named ``__loop_gate`` within the group.
9494

95-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
95+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
9696
:start-after: [START fixed_loop]
9797
:end-before: [END fixed_loop]
9898

9999
For a task-group function with arguments, supply them with ``.partial()`` before
100100
calling ``.loop()``. Use ``.override()`` to configure the group, for example to
101101
give another loop a different ``group_id``:
102102

103-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
103+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
104104
:start-after: [START partial_override_loop]
105105
:end-before: [END partial_override_loop]
106106

@@ -201,7 +201,7 @@ created.
201201
In this example, ``choose_items`` returns the values to map over: ``[1, 2]`` in the
202202
first iteration, then each previous result plus one.
203203

204-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
204+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
205205
:start-after: [START mapped_loop]
206206
:end-before: [END mapped_loop]
207207

‎airflow-core/src/airflow/api_fastapi/execution_api/datamodels/taskinstance.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -421,12 +421,24 @@ def safe_extract_from_orm(cls, data: Any) -> Any:
421421
return values
422422

423423

424+
class LoopContext(BaseModel):
425+
"""Pinned loop definition and enclosing iteration for a task execution."""
426+
427+
node_id: str
428+
index: int = Field(ge=0)
429+
max_iterations: int = Field(gt=0)
430+
terminal_task_id: str
431+
terminal_is_mapped: bool
432+
433+
424434
class TIRunContext(BaseModel):
425435
"""Response schema for TaskInstance run context."""
426436

427437
dag_run: DagRun
428438
"""DAG run information for the task instance."""
429439

440+
loop: LoopContext | None = None
441+
430442
task_reschedule_count: int = 0
431443
"""How many times the task has been rescheduled."""
432444

‎airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py‎

Lines changed: 94 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@
3535
from opentelemetry.trace import StatusCode
3636
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
3737
from pydantic import JsonValue, ValidationError
38-
from sqlalchemy import and_, func, or_, tuple_, update
38+
from sqlalchemy import and_, func, or_, tuple_, union, update
3939
from sqlalchemy.engine import CursorResult
4040
from sqlalchemy.exc import DataError, NoResultFound, SQLAlchemyError
4141
from sqlalchemy.orm import contains_eager, joinedload
@@ -57,6 +57,7 @@
5757
from airflow.api_fastapi.execution_api.datamodels.taskinstance import (
5858
DagRunNoteUpdatePayload,
5959
InactiveAssetsResponse,
60+
LoopContext,
6061
PreviousTIResponse,
6162
PrevSuccessfulDagRunResponse,
6263
TaskBreadcrumbsResponse,
@@ -94,11 +95,12 @@
9495
from airflow.models.asset import AssetActive
9596
from airflow.models.base import ID_LEN
9697
from airflow.models.dag import DagModel
97-
from airflow.models.dagrun import DagRun as DR
98+
from airflow.models.dagrun import DagRun as DR, InvalidLoopDecision
9899
from airflow.models.dynamic_region import AmbiguousProducerError
99100
from airflow.models.hitl import HITLDetail
100101
from airflow.models.log import Log
101102
from airflow.models.task_coordinates import (
103+
LOOP_GATE_OPERATOR,
102104
TaskCoordinateResolver,
103105
build_coordinate_filters,
104106
get_public_region,
@@ -112,7 +114,7 @@
112114
from airflow.state import get_state_backend
113115
from airflow.triggers.base import TriggerEvent
114116
from airflow.utils.sqlalchemy import get_dialect_name
115-
from airflow.utils.state import DagRunState, TaskInstanceState, TerminalTIState
117+
from airflow.utils.state import DagRunState, IntermediateTIState, TaskInstanceState, TerminalTIState
116118

117119
router = VersionedAPIRouter()
118120

@@ -176,6 +178,7 @@ def ti_run(
176178
TI.dag_id,
177179
TI.run_id,
178180
TI.task_id,
181+
TI.region_id,
179182
TI.region_index,
180183
TI.try_number,
181184
TI.max_tries,
@@ -335,6 +338,19 @@ def ti_run(
335338
should_retry=_is_eligible_to_retry(previous_state, ti.try_number, ti.max_tries),
336339
multi_team=conf.getboolean("core", "multi_team"),
337340
)
341+
resolver = TaskCoordinateResolver(dag_bag, session)
342+
if loop_context := resolver.loop_context(ti):
343+
group, index = loop_context
344+
terminal = resolver.get_task(
345+
ti.dag_id, ti.run_id, group.terminal_task_id, dag_version_id=ti.dag_version_id
346+
)
347+
context.loop = LoopContext(
348+
node_id=group.node_id,
349+
index=index,
350+
max_iterations=group.max_iterations,
351+
terminal_task_id=group.terminal_task_id,
352+
terminal_is_mapped=terminal.get_needs_expansion(),
353+
)
338354

339355
# Only set for lang-SDK (foreign-runtime) tasks with a captured TaskFlow arg
340356
# spec; the route excludes unset fields, keeping regular responses lean.
@@ -434,6 +450,35 @@ def ti_update_state(
434450
raise HTTPException(status_code=409, detail={"reason": "invalid_state"})
435451
return Response(status_code=status.HTTP_204_NO_CONTENT)
436452

453+
loop_group = None
454+
loop_gate = None
455+
if isinstance(ti_patch_payload, (TISuccessStatePayload, TITerminalStatePayload)):
456+
gate_run = session.execute(
457+
select(TI.dag_id, TI.run_id).where(
458+
TI.id == task_instance_id,
459+
TI.working_set.is_(True),
460+
TI.operator == LOOP_GATE_OPERATOR,
461+
)
462+
).one_or_none()
463+
if gate_run is not None:
464+
session.execute(
465+
select(DR).where(DR.dag_id == gate_run.dag_id, DR.run_id == gate_run.run_id).with_for_update()
466+
).scalar_one()
467+
loop_gate = session.scalar(
468+
select(TI)
469+
.where(TI.id == task_instance_id, TI.working_set.is_(True))
470+
.with_for_update(of=TI)
471+
.execution_options(populate_existing=True)
472+
)
473+
if loop_gate is not None:
474+
loop_context = TaskCoordinateResolver(dag_bag, session).loop_context(loop_gate)
475+
if loop_context is None:
476+
loop_gate = None
477+
elif loop_context[0].gate_task_id != loop_gate.task_id:
478+
raise HTTPException(status_code=409, detail={"reason": "invalid_loop_gate"})
479+
else:
480+
loop_group = loop_context[0]
481+
437482
old = (
438483
select(
439484
TI.state,
@@ -521,6 +566,25 @@ def ti_update_state(
521566
detail={"reason": "invalid_partition_key", "message": str(e)},
522567
) from e
523568

569+
gate_completed = False
570+
if (
571+
loop_gate is not None
572+
and loop_group is not None
573+
and isinstance(ti_patch_payload, TISuccessStatePayload)
574+
):
575+
try:
576+
loop_gate.dag_run.complete_loop_gate(
577+
loop_gate, loop_group, TaskInstanceState.SUCCESS, session=session
578+
)
579+
gate_completed = True
580+
except InvalidLoopDecision as error:
581+
log.warning("Loop gate success rejected", error=str(error))
582+
ti_patch_payload = _build_rejected_gate_payload(
583+
ti_patch_payload,
584+
reason=f"Loop gate success rejected: {error}",
585+
retry=_is_eligible_to_retry(previous_state, try_number, max_tries),
586+
)
587+
524588
# We exclude_unset to avoid updating fields that are not set in the payload
525589
data = ti_patch_payload.model_dump(
526590
exclude={"task_outlets", "outlet_events", "retry_delay_seconds", "retry_reason"},
@@ -544,6 +608,8 @@ def ti_update_state(
544608
# Let DataErrorHandler return a 422 instead of silently marking the TI FAILED below.
545609
raise
546610
except Exception:
611+
if loop_gate is not None:
612+
raise
547613
# Set a task to failed in case any unexpected exception happened during task state update
548614
log.exception(
549615
"Error updating Task Instance state. Setting the task to failed.",
@@ -594,6 +660,9 @@ def ti_update_state(
594660
# Defer to app-level SQLAlchemyError handler (returns HTTP 500).
595661
raise
596662

663+
if loop_gate is not None and loop_group is not None and not gate_completed:
664+
loop_gate.dag_run.complete_loop_gate(loop_gate, loop_group, updated_state, session=session)
665+
597666
if updated_state == TaskInstanceState.SUCCESS:
598667
if conf.getboolean("state_store", "clear_on_success"):
599668
scope = TaskScope(
@@ -627,6 +696,15 @@ def ti_update_state(
627696
callback()
628697

629698

699+
def _build_rejected_gate_payload(
700+
payload: TISuccessStatePayload, *, reason: str, retry: bool
701+
) -> TIRetryStatePayload | TITerminalStatePayload:
702+
carried = payload.model_dump(include={"end_date", "rendered_map_index"}, exclude_unset=True)
703+
if retry:
704+
return TIRetryStatePayload(state=IntermediateTIState.UP_FOR_RETRY, retry_reason=reason, **carried)
705+
return TITerminalStatePayload(state=TerminalStateNonSuccess.FAILED, retry_reason=reason, **carried)
706+
707+
630708
def _emit_task_span(ti, state, *, resolver: TaskCoordinateResolver):
631709
# just to be safe
632710
if not ti.dag_run:
@@ -939,29 +1017,27 @@ def ti_skip_downstream(
9391017
for task_id, index in targets
9401018
if not resolver.has_regions(dag_id, run_id, task_id)
9411019
]
942-
selected_ids = {
943-
ti.id
1020+
selects = [
1021+
resolver.select_skip_target_ids(caller=caller, task_id=task_id, map_indexes=index)
9441022
for task_id, index in targets
9451023
if (task_id, index) not in plain
946-
for ti in resolver.resolve(
947-
dag_id=dag_id, run_id=run_id, task_id=task_id, caller=caller, map_indexes=index
948-
)
949-
}
1024+
]
9501025
if plain:
951-
selected_ids.update(
952-
session.scalars(
953-
resolver.select_legacy_task_ids(
954-
dag_id=dag_id,
955-
run_id=run_id,
956-
task_ids=[task_id for task_id, index in plain if index is None],
957-
slots=[(task_id, index) for task_id, index in plain if index is not None],
958-
)
1026+
selects.append(
1027+
resolver.select_legacy_task_ids(
1028+
dag_id=dag_id,
1029+
run_id=run_id,
1030+
task_ids=[task_id for task_id, index in plain if index is None],
1031+
slots=[(task_id, index) for task_id, index in plain if index is not None],
9591032
)
9601033
)
9611034
except AmbiguousProducerError as error:
9621035
raise HTTPException(status.HTTP_409_CONFLICT, str(error)) from error
9631036
except ValueError as error:
9641037
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(error)) from error
1038+
if not selects:
1039+
return
1040+
selected_ids = set(session.scalars(union(*selects)))
9651041

9661042
# Don't overwrite tasks that are already executing or finished.
9671043
# See: https://github.com/apache/airflow/issues/59378
@@ -979,6 +1055,7 @@ def ti_skip_downstream(
9791055
query = (
9801056
update(TI)
9811057
.where(
1058+
TI.working_set.is_(True),
9821059
TI.id.in_(selected_ids),
9831060
skippable_state_clause,
9841061
)

0 commit comments

Comments
 (0)