Skip to content

Commit 15a1960

Browse files
committed
fixup! Run loop gates and advance iterations atomically
1 parent 3bdcbd2 commit 15a1960

15 files changed

Lines changed: 311 additions & 85 deletions

File tree

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

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@ count is reached.
5858
This example improves an estimate of the square root of two until the error
5959
is small enough:
6060

61-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
61+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
6262
:start-after: [START refine_estimate]
6363
:end-before: [END refine_estimate]
6464

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

84-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
84+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
8585
:start-after: [START fixed_loop]
8686
:end-before: [END fixed_loop]
8787

8888
For a task-group function with arguments, supply them with ``.partial()`` before
8989
calling ``.loop()``. Use ``.override()`` to configure the group, for example to
9090
give another loop a different ``group_id``:
9191

92-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
92+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
9393
:start-after: [START partial_override_loop]
9494
:end-before: [END partial_override_loop]
9595

@@ -187,7 +187,7 @@ single value or structure. A zero-length expansion skips the mapped task and,
187187
under the gate's ``all_success`` rule, skips the gate; no next iteration is
188188
created.
189189

190-
.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
190+
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
191191
:start-after: [START mapped_loop]
192192
:end-before: [END mapped_loop]
193193

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

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -99,7 +99,11 @@
9999
from airflow.models.dynamic_region import SENTINEL_REGION_ID, AmbiguousProducerError
100100
from airflow.models.hitl import HITLDetail
101101
from airflow.models.log import Log
102-
from airflow.models.task_coordinates import TaskCoordinateResolver, public_map_index_expression
102+
from airflow.models.task_coordinates import (
103+
LOOP_GATE_OPERATOR,
104+
TaskCoordinateResolver,
105+
public_map_index_expression,
106+
)
103107
from airflow.models.taskinstance import TaskInstance as TI, _stop_remaining_tasks
104108
from airflow.models.taskreschedule import TaskReschedule
105109
from airflow.models.trigger import Trigger, handle_event_submit
@@ -172,6 +176,7 @@ def ti_run(
172176
TI.dag_id,
173177
TI.run_id,
174178
TI.task_id,
179+
TI.region_id,
175180
TI.region_index,
176181
TI.try_number,
177182
TI.max_tries,
@@ -450,7 +455,7 @@ def ti_update_state(
450455
select(TI.dag_id, TI.run_id).where(
451456
TI.id == task_instance_id,
452457
TI.working_set.is_(True),
453-
TI.operator == "LoopGateOperator",
458+
TI.operator == LOOP_GATE_OPERATOR,
454459
)
455460
).one_or_none()
456461
if gate_run is not None:
@@ -465,9 +470,12 @@ def ti_update_state(
465470
)
466471
if loop_gate is not None:
467472
loop_context = TaskCoordinateResolver(dag_bag, session).loop_context(loop_gate)
468-
if loop_context is None or loop_context[0].gate_task_id != loop_gate.task_id:
473+
if loop_context is None:
474+
loop_gate = None
475+
elif loop_context[0].gate_task_id != loop_gate.task_id:
469476
raise HTTPException(status_code=409, detail={"reason": "invalid_loop_gate"})
470-
loop_group = loop_context[0]
477+
else:
478+
loop_group = loop_context[0]
471479

472480
old = (
473481
select(

‎airflow-core/src/airflow/api_fastapi/execution_api/versions/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,7 @@
6464
IdentifyArchivedTaskStateUpdates,
6565
)
6666
from airflow.api_fastapi.execution_api.versions.v2026_10_30_xcom_params import (
67+
AddPreviousIterationToXComFilterParams,
6768
AddRegionSelectorsToXComFilterParams,
6869
)
6970

@@ -77,6 +78,7 @@
7778
AddLoopContext,
7879
AddTerminalStateRetryReasonField,
7980
AddMultiTeamToTIRunContext,
81+
AddPreviousIterationToXComFilterParams,
8082
AddStoppedTaskReport,
8183
AddTaskInstanceRegionCoordinates,
8284
AddRegionSelectorsToXComFilterParams,

‎airflow-core/src/airflow/api_fastapi/execution_api/versions/v2026_10_30_xcom_params.py‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,3 +35,14 @@ class AddRegionSelectorsToXComFilterParams(VersionChange):
3535
schema(GetXcomFilterParams).field("region_id").didnt_exist,
3636
schema(GetXcomFilterParams).field("region_index").didnt_exist,
3737
)
38+
39+
40+
class AddPreviousIterationToXComFilterParams(VersionChange):
41+
"""Add the `previous_iteration` field to GetXComSliceFilterParams and GetXcomFilterParams."""
42+
43+
description = __doc__
44+
45+
instructions_to_migrate_to_previous_version = (
46+
schema(GetXComSliceFilterParams).field("previous_iteration").didnt_exist,
47+
schema(GetXcomFilterParams).field("previous_iteration").didnt_exist,
48+
)

airflow-core/docs/authoring-and-scheduling/examples/example_task_loops.py renamed to airflow-core/src/airflow/example_dags/example_task_loops.py

File renamed without changes.

‎airflow-core/src/airflow/jobs/scheduler_job_runner.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -114,7 +114,7 @@
114114
from airflow.models.log import resolve_team_name
115115
from airflow.models.pool import normalize_pool_name_for_stats
116116
from airflow.models.serialized_dag import SerializedDagModel
117-
from airflow.models.task_coordinates import TaskCoordinateResolver
117+
from airflow.models.task_coordinates import LOOP_GATE_OPERATOR, TaskCoordinateResolver
118118
from airflow.models.taskinstance import TaskInstance
119119
from airflow.models.team import Team
120120
from airflow.models.trigger import TRIGGER_FAIL_REPR, Trigger, TriggerFailureReason, handle_event_submit
@@ -1679,7 +1679,8 @@ def process_executor_events(
16791679

16801680
task = dag.get_task(ti.task_id)
16811681
except Exception:
1682-
if state == TaskInstanceState.SUCCESS and ti.operator == "LoopGateOperator":
1682+
if state == TaskInstanceState.SUCCESS and ti.operator == LOOP_GATE_OPERATOR:
1683+
cls.logger().exception("Failing loop gate %s: %s", ti, msg)
16831684
ti.task = None
16841685
ti.handle_failure(error=msg, session=session)
16851686
continue

‎airflow-core/src/airflow/models/task_coordinates.py‎

Lines changed: 50 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
# under the License.
1717
from __future__ import annotations
1818

19-
from typing import TYPE_CHECKING, Any, Protocol
19+
from typing import TYPE_CHECKING, Protocol
2020
from uuid import UUID
2121

2222
import attrs
@@ -49,6 +49,22 @@
4949
from airflow.serialization.definitions.dag import SerializedDAG, SerializedOperator
5050

5151

52+
LOOP_GATE_OPERATOR = "LoopGateOperator"
53+
"""Operator class name the Task SDK gives the gate task of a loop."""
54+
55+
56+
@attrs.frozen(kw_only=True)
57+
class _ProducerRequest:
58+
dag_id: str
59+
run_id: str
60+
task_id: str
61+
is_mapped: bool
62+
context: ProducerContext | None
63+
map_indexes: int | Collection[int] | None
64+
region_id: UUID | None
65+
region_index: int | None
66+
67+
5268
class TaskCoordinate(Protocol):
5369
"""Stored task identity shared by live and retained task data."""
5470

@@ -249,7 +265,12 @@ def has_regions(self, dag_id: str, run_id: str | None, task_id: str) -> bool:
249265
return bool(self.session.scalar(select(query.exists())))
250266

251267
def resolve_dependency(self, caller: TaskInstance, task_id: str) -> tuple[TaskInstance, ...]:
252-
producer = self.get_task(caller.dag_id, caller.run_id, task_id, dag_version_id=caller.dag_version_id)
268+
try:
269+
producer = self.get_task(
270+
caller.dag_id, caller.run_id, task_id, dag_version_id=caller.dag_version_id
271+
)
272+
except TaskNotFound:
273+
return self.resolve(dag_id=caller.dag_id, run_id=caller.run_id, task_id=task_id, caller=caller)
253274
loop = enclosing_loop(producer)
254275
caller_task = self.get_task(
255276
caller.dag_id, caller.run_id, caller.task_id, dag_version_id=caller.dag_version_id
@@ -278,6 +299,18 @@ def resolve_dependency(self, caller: TaskInstance, task_id: str) -> tuple[TaskIn
278299
return tuple(gates[:1])
279300
return self.resolve(dag_id=caller.dag_id, run_id=caller.run_id, task_id=task_id, caller=caller)
280301

302+
@staticmethod
303+
def _filter_region_indexes(
304+
query: Select, region_index: int | None, map_indexes: int | Collection[int] | None
305+
) -> Select:
306+
if region_index is not None:
307+
query = query.where(TaskInstance.region_index == region_index)
308+
if isinstance(map_indexes, int):
309+
query = query.where(TaskInstance.region_index == map_indexes)
310+
elif map_indexes is not None:
311+
query = query.where(TaskInstance.region_index.in_(map_indexes))
312+
return query
313+
281314
def _filter_legacy_producers(
282315
self,
283316
query: Select,
@@ -295,13 +328,7 @@ def _filter_legacy_producers(
295328
TaskInstance.task_id == task_id,
296329
TaskInstance.region_id == SENTINEL_REGION_ID,
297330
)
298-
if region_index is not None:
299-
query = query.where(TaskInstance.region_index == region_index)
300-
if isinstance(map_indexes, int):
301-
query = query.where(TaskInstance.region_index == map_indexes)
302-
elif map_indexes is not None:
303-
query = query.where(TaskInstance.region_index.in_(map_indexes))
304-
return query
331+
return self._filter_region_indexes(query, region_index, map_indexes)
305332

306333
def _filter_removed_task_producers(
307334
self,
@@ -324,13 +351,7 @@ def _filter_removed_task_producers(
324351
)
325352
if region_id is not None:
326353
query = query.where(TaskInstance.region_id == region_id)
327-
if region_index is not None:
328-
query = query.where(TaskInstance.region_index == region_index)
329-
if isinstance(map_indexes, int):
330-
query = query.where(TaskInstance.region_index == map_indexes)
331-
elif map_indexes is not None:
332-
query = query.where(TaskInstance.region_index.in_(map_indexes))
333-
return query
354+
return self._filter_region_indexes(query, region_index, map_indexes)
334355

335356
def _is_legacy_lookup(
336357
self, dag_id: str, run_id: str, task_id: str, region_id: UUID | None, previous_iteration: bool
@@ -351,7 +372,7 @@ def _build_producer_request(
351372
region_index: int | None,
352373
map_indexes: int | Collection[int] | None,
353374
previous_iteration: bool,
354-
) -> dict[str, Any] | None:
375+
) -> _ProducerRequest | None:
355376
shared_run = caller is not None and (caller.dag_id, caller.run_id) == (dag_id, run_id)
356377
try:
357378
task = (
@@ -378,16 +399,16 @@ def _build_producer_request(
378399
)
379400
if previous_iteration and context is None:
380401
raise ValueError("Previous-iteration lookup requires a shared loop scope")
381-
return {
382-
"dag_id": dag_id,
383-
"run_id": run_id,
384-
"task_id": task_id,
385-
"is_mapped": task.get_needs_expansion(),
386-
"context": context,
387-
"map_indexes": map_indexes,
388-
"region_id": region_id,
389-
"region_index": region_index,
390-
}
402+
return _ProducerRequest(
403+
dag_id=dag_id,
404+
run_id=run_id,
405+
task_id=task_id,
406+
is_mapped=task.get_needs_expansion(),
407+
context=context,
408+
map_indexes=map_indexes,
409+
region_id=region_id,
410+
region_index=region_index,
411+
)
391412

392413
def resolve(
393414
self,
@@ -437,7 +458,7 @@ def resolve(
437458
).order_by(TaskInstance.region_index)
438459
)
439460
)
440-
return resolve_current_producers(**request, session=self.session)
461+
return resolve_current_producers(**attrs.asdict(request, recurse=False), session=self.session)
441462

442463
def select_skip_target_ids(
443464
self, *, caller: TaskInstance, task_id: str, map_indexes: int | None = None
@@ -520,4 +541,4 @@ def select_producer_ids(
520541
region_index=region_index,
521542
map_indexes=map_indexes,
522543
)
523-
return select_current_producer_ids(**request, session=self.session)
544+
return select_current_producer_ids(**attrs.asdict(request, recurse=False), session=self.session)

‎airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py‎

Lines changed: 42 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,7 @@
5151
from airflow.api_fastapi.execution_api.app import lifespan
5252
from airflow.api_fastapi.execution_api.datamodels.taskinstance import TISuccessStatePayload
5353
from airflow.api_fastapi.execution_api.datamodels.token import TIClaims, TIToken
54+
from airflow.api_fastapi.execution_api.routes import task_instances as task_instances_route
5455
from airflow.api_fastapi.execution_api.routes.task_instances import _emit_task_span, ti_update_state
5556
from airflow.api_fastapi.execution_api.routes.xcoms import set_xcom
5657
from airflow.api_fastapi.execution_api.security import require_auth
@@ -62,7 +63,7 @@
6263
from airflow.models.dagbag import DBDagBag
6364
from airflow.models.dynamic_region import DynamicRegion
6465
from airflow.models.log import Log
65-
from airflow.models.task_coordinates import TaskCoordinateResolver
66+
from airflow.models.task_coordinates import LOOP_GATE_OPERATOR, TaskCoordinateResolver
6667
from airflow.models.task_state_store import TaskStateStoreModel
6768
from airflow.models.taskinstance import TaskInstance, clear_task_instances
6869
from airflow.models.xcom import XComModel, XComModelV2
@@ -1798,6 +1799,7 @@ def complete():
17981799
) == [0, 1]
17991800
assert session.scalar(select(XComModel).where(XComModel.task_id == gate.task_id)) is None
18001801

1802+
@pytest.mark.backend("mysql", "postgres")
18011803
def test_loop_decision_rewrite_and_completion_use_same_lock_order(
18021804
self, session, running_loop_gate, mocker
18031805
):
@@ -1905,6 +1907,45 @@ def test_late_loop_decision_cannot_recreate_consumed_signal(self, client, sessio
19051907
assert response.status_code == 409
19061908
assert session.scalar(select(XComModel).where(XComModel.task_id == gate.task_id)) is None
19071909

1910+
def test_loop_gate_state_update_error_rolls_back_next_pass_and_keeps_decision(
1911+
self, client, session, running_loop_gate, mocker
1912+
):
1913+
gate = running_loop_gate
1914+
mocker.patch.object(
1915+
task_instances_route,
1916+
"_create_ti_state_update_query_and_update_state",
1917+
autospec=True,
1918+
side_effect=StaleDataError("injected state update failure"),
1919+
)
1920+
1921+
response = client.patch(
1922+
f"/execution/task-instances/{gate.id}/state",
1923+
json={"state": "success", "end_date": DEFAULT_END_DATE.isoformat()},
1924+
)
1925+
1926+
assert response.status_code == 500
1927+
session.expire_all()
1928+
assert gate.state == State.RUNNING
1929+
assert len(gate.dag_run.get_task_instances(session=session)) == 2
1930+
assert session.scalar(select(XComModel.value).where(XComModel.task_id == gate.task_id)) == "continue"
1931+
1932+
@pytest.mark.parametrize("state", ["success", "failed"])
1933+
def test_non_loop_operator_named_loop_gate_completes_normally(
1934+
self, client, session, create_task_instance, state
1935+
):
1936+
ti = create_task_instance(task_id="user_gate", start_date=DEFAULT_START_DATE, state=State.RUNNING)
1937+
ti.operator = LOOP_GATE_OPERATOR
1938+
session.commit()
1939+
1940+
response = client.patch(
1941+
f"/execution/task-instances/{ti.id}/state",
1942+
json={"state": state, "end_date": DEFAULT_END_DATE.isoformat()},
1943+
)
1944+
1945+
assert response.status_code == 204
1946+
session.expire_all()
1947+
assert ti.state == state
1948+
19081949
def test_loop_gate_inner_insert_error_is_not_swallowed(self, client, session, running_loop_gate, mocker):
19091950
gate = running_loop_gate
19101951
bulk_insert = mocker.patch.object(

‎airflow-core/tests/unit/jobs/test_scheduler_job.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1174,7 +1174,11 @@ def body():
11741174
assert (gate.id != original_id) is bool(retries)
11751175
callback.assert_not_called()
11761176
if retries:
1177-
archived = session.scalar(select(TaskInstance).where(TaskInstance.id == original_id))
1177+
archived = session.scalar(
1178+
select(TaskInstance)
1179+
.where(TaskInstance.id == original_id)
1180+
.execution_options(include_all_attempts=True)
1181+
)
11781182
assert archived.working_set is None
11791183
assert len(dr.get_task_instances(session=session)) == 2
11801184

0 commit comments

Comments
 (0)