kaxil commented on code in PR #74348:
URL: https://github.com/apache/airflow/pull/74348#discussion_r4222281427


##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py:
##########
@@ -568,6 +592,42 @@ def set_xcom(
     """Set an Airflow XCom."""
     from airflow.configuration import conf
 
+    region_id: UUID | None = None
+    region_index: int | None = None
+    if loop_decision:
+        caller = session.get(TaskInstance, token.id)
+        if (
+            key != LOOP_DECISION_KEY
+            or value not in ("continue", "stop")
+            or mapped_length is not None
+            or dag_result
+            or map_index != -1
+            or caller is None
+            or (dag_id, run_id, task_id) != (caller.dag_id, caller.run_id, 
caller.task_id)
+        ):
+            raise HTTPException(status.HTTP_400_BAD_REQUEST, "Invalid loop 
decision write")
+        session.execute(
+            select(DagRun)
+            .where(DagRun.dag_id == caller.dag_id, DagRun.run_id == 
caller.run_id)
+            .with_for_update()
+        ).scalar_one()
+        caller = session.scalar(
+            select(TaskInstance)
+            .where(TaskInstance.id == token.id)
+            .with_for_update()
+            .execution_options(populate_existing=True)
+        )
+        if caller is None:
+            raise HTTPException(status.HTTP_409_CONFLICT, "Loop gate is no 
longer running")
+        loop = TaskCoordinateResolver(dag_bag, session).loop_context(caller)

Review Comment:
   A gate whose decision fails `complete_loop_gate`'s checks still has its 
success rewritten to up_for_retry or failed after the worker has already run 
`on_success_callback`, so its failure and retry callbacks never fire. 
`set_xcom` already loads the pinned loop here; running the same `at_limit` and 
fixed-count checks before writing the decision would fail the gate inside 
`execute()`, where normal retry and failure handling applies, and leave the 
rewrite in `ti_update_state` as a backstop.



##########
devel-common/src/tests_common/test_utils/asserts.py:
##########
@@ -282,3 +284,18 @@ def capture(orm_execute_state: ORMExecuteState) -> None:
                 for element in joined - set(linter.froms)
             )
     assert not problems, "Cartesian product in generated SQL:\n" + 
"\n".join(sorted(set(problems)))
+
+
+@contextmanager
+def count_loaded_task_instances(task_id: str) -> Generator[list[TaskInstance], 
None, None]:
+    """Collect the ``task_id`` task instance rows the ORM loads inside the 
block."""
+    from airflow.models.taskinstance import TaskInstance
+
+    loaded: list[TaskInstance] = []
+    listener = loaded.append

Review Comment:
   SQLAlchemy calls a `load` listener with `(target, context)`, so 
`loaded.append` raises `TypeError: list.append() takes exactly one argument (2 
given)` on the first TaskInstance the ORM loads, and `loaded` never collects 
anything. The two callers only pass because nothing loads. Something like `def 
listener(target, _context): if target.task_id == task_id: 
loaded.append(target)` would make the helper report the rows instead of failing 
from inside the ORM.



##########
task-sdk/tests/task_sdk/execution_time/test_loop.py:
##########
@@ -0,0 +1,225 @@
+# 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.
+from __future__ import annotations
+
+import pytest
+
+from airflow.sdk import DAG, BaseOperator, TaskInstanceState, task, task_group
+from airflow.sdk.api.datamodels._generated import LoopContext
+from airflow.sdk.bases.xcom import BaseXCom
+from airflow.sdk.definitions._internal.loop import create_loop
+from airflow.sdk.exceptions import AirflowFailException
+from airflow.sdk.execution_time import task_runner
+from airflow.sdk.execution_time.comms import (
+    GetXCom,
+    GetXComCount,
+    GetXComSequenceItem,
+    GetXComSequenceSlice,
+    SetXCom,
+    XComCountResponse,
+    XComResult,
+    XComSequenceIndexResult,
+    XComSequenceSliceResult,
+)
+from airflow.sdk.execution_time.lazy_sequence import LazyXComSequence
+
+
[email protected]
+def loop_ti(create_runtime_ti):
+    def make(*, index=0, max_iterations=3, until=None, map_index=-1):
+        @task_group
+        def body():
+            BaseOperator(task_id="terminal")
+
+        with DAG("loop_runtime", schedule=None) as dag:
+            group = create_loop(body, max_iterations=max_iterations, 
until=until)
+        ti = create_runtime_ti(task=dag.get_task(group.gate_task_id), 
map_index=map_index)
+        ti._ti_context_from_server.loop = LoopContext(
+            node_id=group.group_id,
+            index=index,
+            max_iterations=max_iterations,
+            terminal_task_id=group.terminal_task_id,
+            terminal_is_mapped=False,
+        )
+        return ti
+
+    return make
+
+
+def test_first_iteration_previous_does_not_read_xcom(loop_ti, 
mock_supervisor_comms):
+    ti = loop_ti(map_index=7)
+    loop = ti.get_template_context()["loop"]
+
+    assert loop.index == 0
+    assert loop.max_iterations == 3
+    assert loop.previous is None
+    assert ti.map_index == 7
+    mock_supervisor_comms.send.assert_not_called()
+
+
[email protected]("mapped", [False, True])
+def 
test_body_callable_receives_loop_context_after_unmapping(create_runtime_ti, 
mapped):
+    @task
+    def terminal(value, *, loop, ti):
+        return value, loop.index, ti.map_index
+
+    @task_group
+    def body():
+        if mapped:
+            terminal.expand(value=[5])
+        else:
+            terminal(5)
+
+    with DAG("loop_body_runtime", schedule=None) as dag:
+        group = create_loop(body, max_iterations=4)
+    operator = dag.get_task(group.terminal_task_id)
+    if mapped:
+        operator = operator.unmap({"op_kwargs": {"value": 5}})
+    ti = create_runtime_ti(task=operator, map_index=0 if mapped else -1)
+    ti._ti_context_from_server.loop = LoopContext(
+        node_id=group.group_id,
+        index=2,
+        max_iterations=4,
+        terminal_task_id=group.terminal_task_id,
+        terminal_is_mapped=mapped,
+    )
+
+    assert operator.execute(ti.get_template_context()) == (5, 2, 0 if mapped 
else -1)
+
+
[email protected]("value", [0, False, [], ""])
+def test_previous_and_current_result_preserve_falsey_values(loop_ti, 
mock_supervisor_comms, value):
+    ti = loop_ti(index=1)
+    mock_supervisor_comms.send.return_value = XComResult(key="return_value", 
value=value)
+    loop = ti.get_template_context()["loop"]
+
+    assert loop.previous == value
+    previous = mock_supervisor_comms.send.call_args.args[0]
+    assert isinstance(previous, GetXCom)
+    assert previous.task_id == "body.terminal"
+    assert previous.previous_iteration is True
+    assert loop.result == value
+    assert mock_supervisor_comms.send.call_args.args[0].previous_iteration is 
False
+
+
[email protected](
+    ("index", "condition", "decision"),
+    [(0, None, "continue"), (2, None, "stop"), (0, False, "continue"), (0, 
True, "stop"), (2, True, "stop")],
+)
+def test_gate_publishes_successful_decision(loop_ti, mock_supervisor_comms, 
index, condition, decision):
+    def until(*, loop):
+        assert loop.index == index
+        return condition
+
+    ti = loop_ti(index=index, until=until if condition is not None else None)
+    context = ti.get_template_context()
+
+    ti.task.execute(context)
+
+    message = mock_supervisor_comms.send.call_args.args[0]
+    assert isinstance(message, SetXCom)
+    assert message.key == "_airflow_loop_decision"
+    assert message.value == decision
+    assert message.loop_decision is True
+
+
[email protected]("raises", [False, True])
+def test_unsuccessful_gate_does_not_publish_decision(loop_ti, 
mock_supervisor_comms, raises):
+    def until(*, loop):
+        if raises:
+            raise ValueError("condition failed")
+        return False
+
+    ti = loop_ti(index=2, until=until)
+    with pytest.raises(Exception, match="condition failed" if raises else 
"max_iterations"):

Review Comment:
   `LoopMaxIterationsExceeded` is still an `AirflowException`, so a gate that 
picks up `retries` from `default_args` retries the cap failure, calling `until` 
again on the same result, before it fails. loops.rst still says the final gate 
fails at the limit. Either `AirflowFailException` or a sentence in the docs 
would settle it, and matching `LoopMaxIterationsExceeded` here instead of 
`Exception` would pin whichever behaviour you pick.



##########
airflow-core/src/airflow/models/dynamic_region.py:
##########
@@ -495,46 +561,85 @@ def resolve_current_producers(
                 position = position[0], position[1] - 1
                 if position[1] < 0:
                     return ()
-    query = select(TaskInstance).where(
-        TaskInstance.dag_id == dag_id,
-        TaskInstance.run_id == run_id,
-        TaskInstance.task_id == task_id,
-        TaskInstance.working_set.is_(True),
-    )
-    if region_id is not None:
-        query = query.where(TaskInstance.region_id == region_id)
-    if region_index is not None:
-        query = query.where(TaskInstance.region_index == region_index)
-    if position is not None:
-        query = 
query.where(_build_loop_pass_filter(_load_loop_pass_regions(position, 
session=session)))
-    candidates = session.scalars(query).all()
-    if region_id is None and position is None:
-        regions.update(
-            load_region_ancestry(
-                {ti.region_id for ti in candidates} - set(regions),
-                dag_id=dag_id,
-                run_id=run_id,
-                session=session,
-            )
+
+    candidates = session.scalars(
+        _filter_producers(
+            select(TaskInstance),
+            dag_id=dag_id,
+            run_id=run_id,
+            task_id=task_id,
+            is_mapped=is_mapped,
+            map_indexes=map_indexes,
+            region_id=region_id,
+            region_index=region_index,
+            top_level_only=region_id is None and position is None,
+            loop_pass=_load_loop_pass_regions(position, session=session) if 
position is not None else None,
         )
+    ).all()
 
     selected: dict[int, TaskInstance] = {}
     for ti in candidates:
-        if region_id is None and position is None:
-            if ti.region_id != SENTINEL_REGION_ID:
-                if not is_mapped or regions[ti.region_id].parent_region_id is 
not None:
-                    continue
-            elif not is_mapped and ti.region_index != -1:
-                continue
         public_index = ti.region_index if is_mapped else -1
-        if isinstance(map_indexes, int):
-            if public_index != map_indexes:
-                continue
-        elif map_indexes is not None and public_index not in map_indexes:
-            continue
         if public_index in selected:
             raise AmbiguousProducerError(
                 f"Multiple live producers for {dag_id}/{run_id}/{task_id} 
index {public_index}"
             )
         selected[public_index] = ti
     return tuple(selected[index] for index in sorted(selected))
+
+
+def select_current_producer_ids(
+    *,
+    dag_id: str,
+    run_id: str,
+    task_id: str,
+    is_mapped: bool,
+    context: ProducerContext | None = None,
+    map_indexes: int | Collection[int] | None = None,
+    region_id: UUID | None = None,
+    region_index: int | None = None,
+    session: Session,
+) -> Select[tuple[UUID]]:
+    """
+    Select the ids of the live producers :func:`resolve_current_producers` 
would return.
+
+    A task whose live executions all share one region cannot have two 
producers at one index, so
+    that common case stays a pure SQL selection whose cost does not depend on 
the number of
+    mapped instances. Everything else resolves the rows and pins their ids.
+    """
+    from airflow.models.taskinstance import TaskInstance
+
+    _validate_producer_request(context=context, region_id=region_id, 
region_index=region_index)
+    filter_producers = partial(
+        _filter_producers,
+        dag_id=dag_id,
+        run_id=run_id,
+        task_id=task_id,
+        is_mapped=is_mapped,
+        map_indexes=map_indexes,
+        region_id=region_id,
+        region_index=region_index,
+        top_level_only=region_id is None,
+    )
+    query = filter_producers(select(TaskInstance.id))
+    if region_id is None:
+        if context is not None:
+            single_region = False
+        else:
+            first_region = 
filter_producers(select(TaskInstance.region_id)).limit(1)
+            other_region = filter_producers(select(TaskInstance.id)).where(
+                TaskInstance.region_id != first_region.scalar_subquery()
+            )
+            single_region = not session.scalar(select(other_region.exists()))
+        if not single_region:
+            producers = resolve_current_producers(

Review Comment:
   Moving the pass filter into SQL fixed the growth per pass. With a context, 
though, this fallback still always runs and loads full `TaskInstance` rows, so 
`min(loop.result)` over M mapped slots is still M+1 item reads that each 
hydrate M rows (about a million row loads at M=1000). Selecting only `id` and 
`region_index` for this path would keep the ambiguity check without hydrating 
the rows.



##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -2177,6 +2186,121 @@ def create_ti(task: Operator, indexes: Iterable[int], 
region_id: UUID) -> Iterat
             creator = create_ti
         return creator
 
+    def complete_loop_gate(
+        self,
+        gate: TI,
+        group: SerializedLoopTaskGroup,
+        state: TaskInstanceState,
+        *,
+        session: Session,
+    ) -> None:
+        """Consume a gate decision while the caller holds the DagRun and TI 
locks."""
+        from airflow.models.xcom import XComModelV2
+        from airflow.settings import task_instance_mutation_hook
+
+        if gate.dag_version_id is None:
+            raise InvalidLoopDecision("Loop gate requires a pinned DAG 
version")
+        signal = XComModelV2.get_for_attempt(gate.id, LOOP_DECISION_KEY, 
session=session)
+        later_gate = session.scalar(
+            select(TI.id)
+            .where(
+                TI.working_set.is_(True),
+                TI.dag_id == self.dag_id,
+                TI.run_id == self.run_id,
+                TI.task_id == gate.task_id,
+                TI.region_id == gate.region_id,
+                TI.region_index > gate.region_index,
+            )
+            .limit(1)
+        )
+        if state != TaskInstanceState.SUCCESS or later_gate:
+            if signal is not None:
+                session.delete(signal)
+            return
+        decision = signal.value if signal is not None else None
+        at_limit = gate.region_index + 1 >= group.max_iterations
+        if decision not in ("continue", "stop"):
+            raise InvalidLoopDecision("Successful loop gate requires a 
decision")
+        if decision == "continue" and at_limit:
+            raise InvalidLoopDecision("Loop cannot continue beyond its 
iteration limit")
+        if not group.has_until and (decision == "stop") != at_limit:
+            raise InvalidLoopDecision("Fixed-count loop decision does not 
match its iteration limit")
+        session.delete(signal)
+        if decision == "stop":
+            return
+        created_counts: dict[str, int] = defaultdict(int)
+        hook_is_noop: Literal[True, False] = 
getattr(task_instance_mutation_hook, "is_noop", False)
+        creator = self._get_task_creator(
+            created_counts, task_instance_mutation_hook, hook_is_noop, 
gate.dag_version_id
+        )
+        tasks = list(
+            self._create_tasks(
+                group.iter_tasks(),
+                creator,
+                session=session,
+                parent_region=(gate.region_id, gate.region_index + 1),
+            )
+        )
+        self._create_task_instances(
+            self.dag_id, tasks, created_counts, hook_is_noop, session=session, 
propagate_errors=True
+        )
+
+    def _create_initial_tasks(
+        self,
+        tasks: Iterable[Operator],
+        task_creator: Callable[[Operator, Iterable[int], UUID], CreatedTasks],
+        *,
+        session: Session,
+    ) -> CreatedTasks:
+        from airflow.models.task_coordinates import enclosing_loop
+
+        ordinary_tasks = []
+        loop_tasks = defaultdict(list)
+        loops = {}
+        for task in tasks:
+            if loop := enclosing_loop(task):
+                loops[loop.group_id] = loop
+                loop_tasks[loop.group_id].append(task)
+            else:
+                ordinary_tasks.append(task)
+        yield from self._create_tasks(ordinary_tasks, task_creator, 
session=session, expand_literals=True)
+        for group_id, members in loop_tasks.items():
+            coordinates = (
+                session.execute(
+                    select(TI.region_id, TI.region_index).where(
+                        TI.working_set.is_(True),
+                        TI.dag_id == self.dag_id,
+                        TI.run_id == self.run_id,
+                        TI.task_id == loops[group_id].gate_task_id,
+                    )
+                )
+                .tuples()
+                .all()
+            )
+            if not coordinates:
+                if session.scalar(
+                    select(DynamicRegion.id)
+                    .where(
+                        DynamicRegion.dag_id == self.dag_id,
+                        DynamicRegion.run_id == self.run_id,
+                        DynamicRegion.node_id == group_id,
+                    )
+                    .limit(1)
+                ):
+                    raise ValueError(f"Loop {group_id!r} has regions but no 
live gate")

Review Comment:
   This still raises when the live gate id no longer matches the loop's 
`gate_task_id`, for example after renaming the `until` function or adding 
`until=` to a fixed-count loop on an unversioned bundle. The scheduler's 
catch-all just logs it, and the loop never gets a gate under the new id. Could 
the live passes come from the loop's regions (`DynamicRegion.node_id == 
group_id`) instead of gate TIs, or could this skip only this loop with a 
warning?



##########
airflow-core/src/airflow/models/task_coordinates.py:
##########
@@ -401,37 +503,137 @@ def resolve(
             map_indexes=map_indexes,
             region_id=region_id,
             region_index=region_index,
-            session=self.session,
         )
 
-    def _resolve_removed_task(
+    def resolve(
         self,
         *,
         dag_id: str,
         run_id: str,
         task_id: str,
-        region_id: UUID | None,
-        region_index: int | None,
-        map_indexes: int | Collection[int] | None,
+        caller: TaskInstance | None = None,
+        region_id: UUID | None = None,
+        region_index: int | None = None,
+        map_indexes: int | Collection[int] | None = None,
+        previous_iteration: bool = False,
     ) -> tuple[TaskInstance, ...]:
-        """Find live rows of a task whose definition is gone, using only its 
own expansion regions."""
-        query = (
-            select(TaskInstance)
-            .join(DynamicRegion, DynamicRegion.id == TaskInstance.region_id)
-            .where(
-                TaskInstance.working_set.is_(True),
-                TaskInstance.dag_id == dag_id,
-                TaskInstance.run_id == run_id,
-                TaskInstance.task_id == task_id,
-                DynamicRegion.node_id == task_id,
+        if region_index is not None and region_id is None:
+            raise ValueError("region_index requires an explicit producer 
region_id")
+        if self._is_legacy_lookup(dag_id, run_id, task_id, region_id, 
previous_iteration):
+            query = self._filter_legacy_producers(
+                select(TaskInstance),
+                dag_id=dag_id,
+                run_id=run_id,
+                task_id=task_id,
+                region_index=region_index,
+                map_indexes=map_indexes,
             )
+            return 
tuple(self.session.scalars(query.order_by(TaskInstance.region_index)))
+        request = self._build_producer_request(
+            dag_id=dag_id,
+            run_id=run_id,
+            task_id=task_id,
+            caller=caller,
+            region_id=region_id,
+            region_index=region_index,
+            map_indexes=map_indexes,
+            previous_iteration=previous_iteration,
         )
-        if region_id is not None:
-            query = query.where(TaskInstance.region_id == region_id)
-        if region_index is not None:
-            query = query.where(TaskInstance.region_index == region_index)
-        if isinstance(map_indexes, int):
-            query = query.where(TaskInstance.region_index == map_indexes)
-        elif map_indexes is not None:
-            query = query.where(TaskInstance.region_index.in_(map_indexes))
-        return 
tuple(self.session.scalars(query.order_by(TaskInstance.region_index)))
+        if request is None:
+            return tuple(
+                self.session.scalars(
+                    self._filter_removed_task_producers(
+                        select(TaskInstance),
+                        dag_id=dag_id,
+                        run_id=run_id,
+                        task_id=task_id,
+                        region_id=region_id,
+                        region_index=region_index,
+                        map_indexes=map_indexes,
+                    ).order_by(TaskInstance.region_index)
+                )
+            )
+        return resolve_current_producers(**attrs.asdict(request, 
recurse=False), session=self.session)

Review Comment:
   `_ProducerRequest` is typed now, but `**attrs.asdict(request, 
recurse=False)` still hides the kwargs from mypy, so a field renamed on either 
side only fails at runtime. Passing the fields explicitly, or having 
`resolve_current_producers` and `select_current_producer_ids` take the request 
itself, would keep the check.



##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py:
##########
@@ -559,6 +582,7 @@ def set_xcom(
         ),
     ] = None,
     map_index: Annotated[int, Query()] = -1,
+    loop_decision: bool = False,

Review Comment:
   The route only accepts `loop_decision=True` with `key == LOOP_DECISION_KEY` 
and rejects that key without it, so the flag never says anything the key 
doesn't. Could the route key off the reserved key and drop the flag? That would 
also keep `loop_decision` out of `SetXCom` in the cross-language supervisor 
schema.



##########
airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py:
##########
@@ -1490,6 +1551,496 @@ def test_ti_run_creates_audit_log(self, client, 
session, create_task_instance, t
 
 
 class TestTIUpdateState:
+    def test_loop_continues_while_earlier_body_branch_is_running(self, client, 
session, dag_maker):
+        @task_group
+        def body():
+            left = PythonOperator(task_id="left", python_callable=list)
+            right = PythonOperator(task_id="right", python_callable=list)
+            terminal = PythonOperator(
+                task_id="terminal", python_callable=list, 
trigger_rule=TriggerRule.ONE_SUCCESS
+            )
+            [left, right] >> terminal
+
+        with dag_maker(serialized=True, session=session):
+            loop = create_loop(body, max_iterations=2)
+        dr = dag_maker.create_dagrun()
+        dr.dag = dag_maker.serialized_dag
+        tis = {ti.task_id: ti for ti in dr.task_instances}
+        tis["body.left"].state = State.SUCCESS
+        right = tis["body.right"]
+        right.state = State.RUNNING
+        right_id = right.id
+        terminal = tis["body.terminal"]
+        gate = tis[loop.gate_task_id]
+        session.flush()
+        assert terminal in 
dr.task_instance_scheduling_decisions(session=session).schedulable_tis
+        terminal.state, terminal.start_date = State.RUNNING, DEFAULT_START_DATE
+        session.commit()
+        assert (
+            client.patch(
+                f"/execution/task-instances/{terminal.id}/state",
+                json={"state": "success", "end_date": 
DEFAULT_END_DATE.isoformat()},
+            ).status_code
+            == 204
+        )
+        session.expire_all()
+        assert gate in 
dr.task_instance_scheduling_decisions(session=session).schedulable_tis
+        gate.state, gate.start_date = State.RUNNING, DEFAULT_START_DATE
+        session.commit()
+        exec_app = client.app.routes[-1].app
+        exec_app.dependency_overrides[require_auth] = lambda: 
TIToken(id=gate.id, claims=TIClaims())
+        assert (
+            client.post(
+                
f"/execution/xcoms/{dr.dag_id}/{dr.run_id}/{gate.task_id}/_airflow_loop_decision",
+                params={"loop_decision": True},
+                json="continue",
+            ).status_code
+            == 201
+        )
+
+        response = client.patch(
+            f"/execution/task-instances/{gate.id}/state",
+            json={"state": "success", "end_date": 
DEFAULT_END_DATE.isoformat()},
+        )
+
+        assert response.status_code == 204, response.text
+        session.expire_all()
+        assert session.get(TaskInstance, right_id).state == State.RUNNING
+        ready = 
dr.task_instance_scheduling_decisions(session=session).schedulable_tis
+        assert {(ti.task_id, ti.region_index) for ti in ready} == 
{("body.left", 1), ("body.right", 1)}
+
+    @conf_vars({("state_store", "clear_on_success"): "true"})
+    @pytest.mark.parametrize("mapped", [False, True])
+    @pytest.mark.parametrize(
+        ("decision", "state", "max_iterations", "conditional", "rejected", 
"passes"),
+        [
+            ("continue", State.SUCCESS, 3, False, False, [0, 1]),
+            ("stop", State.SUCCESS, 3, False, True, [0]),
+            (None, State.SUCCESS, 3, False, True, [0]),
+            ("continue", State.FAILED, 3, False, False, [0]),
+            ("continue", State.SKIPPED, 3, False, False, [0]),
+            ("stop", State.SUCCESS, 1, False, False, [0]),
+            ("continue", State.SUCCESS, 1, False, True, [0]),
+            ("stop", State.SUCCESS, 3, True, False, [0]),
+            ("stop", State.SUCCESS, 1, True, False, [0]),
+            ("continue", State.SUCCESS, 1, True, True, [0]),
+            (None, State.FAILED, 1, True, False, [0]),
+        ],
+    )
+    def test_loop_gate_completion_consumes_decision_atomically(
+        self,
+        client,
+        session,
+        dag_maker,
+        decision,
+        state,
+        max_iterations,
+        conditional,
+        mapped,
+        rejected,
+        passes,
+        mocker,
+    ):
+        backend = mocker.create_autospec(MetastoreBackend, instance=True)
+        mocker.patch(
+            
"airflow.api_fastapi.execution_api.routes.task_instances.get_state_backend",
+            autospec=True,
+            return_value=backend,
+        )
+
+        @task_group
+        def body():
+            if mapped:
+                PythonOperator.partial(task_id="terminal", 
python_callable=list).expand(op_kwargs=[{}, {}])
+            else:
+                PythonOperator(task_id="terminal", python_callable=list)
+
+        with dag_maker(serialized=True, session=session):
+            loop = create_loop(
+                body, max_iterations=max_iterations, until=(lambda loop: True) 
if conditional else None
+            )
+        dr = dag_maker.create_dagrun()
+        gate = next(ti for ti in dr.task_instances if ti.task_id == 
loop.gate_task_id)
+        gate.state = State.RUNNING
+        gate.start_date = DEFAULT_START_DATE
+        session.commit()
+        if decision:
+            exec_app = client.app.routes[-1].app
+            exec_app.dependency_overrides[require_auth] = lambda: 
TIToken(id=gate.id, claims=TIClaims())
+            response = client.post(
+                
f"/execution/xcoms/{dr.dag_id}/{dr.run_id}/{gate.task_id}/_airflow_loop_decision",
+                params={"loop_decision": True},
+                json=decision,
+            )
+            assert response.status_code == 201
+
+        response = client.patch(
+            f"/execution/task-instances/{gate.id}/state",
+            json={"state": state, "end_date": DEFAULT_END_DATE.isoformat()},
+        )
+
+        assert response.status_code == 204, response.text
+        assert backend.clear.call_count == (state == State.SUCCESS and not 
rejected)
+        session.expire_all()
+        assert gate.state == (State.FAILED if rejected else state)
+        current = dr.get_task_instances(session=session)
+        assert sorted(ti.region_index for ti in current if ti.task_id == 
gate.task_id) == passes
+        if mapped:
+            children = session.scalars(
+                select(DynamicRegion).where(DynamicRegion.parent_region_id == 
gate.region_id)
+            ).all()
+            assert sorted(region.parent_region_index for region in children) 
== passes
+            region_ids = {gate.region_id, *(region.id for region in children)}
+            assert all(ti.region_id in region_ids for ti in current)
+        else:
+            assert all(ti.region_id == gate.region_id for ti in current)
+        signal = session.scalar(select(XComModel).where(XComModel.task_id == 
gate.task_id))
+        assert signal is None
+        if not rejected:
+            duplicate = client.patch(
+                f"/execution/task-instances/{gate.id}/state",
+                json={"state": state, "end_date": 
DEFAULT_END_DATE.isoformat()},
+            )
+            assert duplicate.status_code == 200
+            assert (
+                sorted(
+                    ti.region_index
+                    for ti in dr.get_task_instances(session=session)
+                    if ti.task_id == gate.task_id
+                )
+                == passes
+            )
+        if mapped and passes == [0, 1]:
+            successor_region = next(region for region in children if 
region.parent_region_index == 1)
+            assert [ti.region_index for ti in current if ti.region_id == 
successor_region.id] == [-1]
+            dr.dag = dag_maker.serialized_dag
+
+            dr.task_instance_scheduling_decisions(session=session)
+
+            assert sorted(
+                ti.region_index
+                for ti in dr.get_task_instances(session=session)
+                if ti.region_id == successor_region.id
+            ) == [0, 1]
+
+    @pytest.fixture
+    def running_loop_gate(self, session, dag_maker):
+        @task_group
+        def body():
+            PythonOperator(task_id="terminal", python_callable=list)
+
+        with dag_maker(serialized=True, session=session):
+            loop = create_loop(body, max_iterations=3)
+        dr = dag_maker.create_dagrun()
+        gate = next(ti for ti in dr.task_instances if ti.task_id == 
loop.gate_task_id)
+        gate.state, gate.start_date = State.RUNNING, DEFAULT_START_DATE
+        XComModel.set_for_attempt(
+            task_instance_id=gate.id,
+            key="_airflow_loop_decision",
+            value="continue",
+            serialize=False,
+            session=session,
+        )
+        session.commit()
+        return gate
+
+    def 
test_loop_gate_materialization_error_rolls_back_state_signal_and_successor(
+        self, client, session, running_loop_gate, mocker
+    ):
+        gate = running_loop_gate
+        dr = gate.dag_run
+        original = DagRun._create_task_instances
+        observed = []
+
+        def fail_after_insert(*args, **kwargs):
+            original(*args, **kwargs)
+            with Session(bind=session.get_bind()) as observer:
+                observed.append(observer.get(TaskInstance, gate.id).state)
+                observed.append(
+                    observer.scalars(
+                        select(TaskInstance.region_index).where(
+                            TaskInstance.dag_id == dr.dag_id, 
TaskInstance.task_id == gate.task_id
+                        )
+                    ).all()
+                )
+            raise StaleDataError("injected materialization failure")
+
+        mocker.patch.object(DagRun, "_create_task_instances", autospec=True, 
side_effect=fail_after_insert)
+        response = client.patch(
+            f"/execution/task-instances/{gate.id}/state",
+            json={"state": "success", "end_date": 
DEFAULT_END_DATE.isoformat()},
+        )
+
+        assert response.status_code == 500
+        session.expire_all()
+        assert observed == [State.RUNNING, [0]]
+        assert gate.state == State.RUNNING
+        assert len(dr.get_task_instances(session=session)) == 2
+        assert session.scalar(select(XComModel.value).where(XComModel.task_id 
== gate.task_id)) == "continue"
+
+    @pytest.mark.parametrize(
+        ("max_tries", "expected_state"),
+        [(0, State.FAILED), (1, State.UP_FOR_RETRY)],
+    )
+    def test_loop_gate_with_invalid_decision_leaves_running_in_same_request(
+        self, client, session, running_loop_gate, max_tries, expected_state
+    ):
+        gate = running_loop_gate
+        dr = gate.dag_run
+        gate_id = gate.id
+        gate.max_tries = max_tries
+        session.execute(delete(XComModelV2).where(XComModelV2.task_instance_id 
== gate.id))
+        session.commit()
+
+        response = client.patch(
+            f"/execution/task-instances/{gate.id}/state",
+            json={"state": "success", "end_date": 
DEFAULT_END_DATE.isoformat()},
+        )
+
+        assert response.status_code == 204
+        session.expire_all()
+        current = dr.get_task_instances(session=session)
+        assert sorted(ti.state for ti in current if ti.task_id == 
gate.task_id) == [expected_state]
+        assert len(current) == 2
+        assert "requires a decision" in session.get(TaskInstance, 
gate_id).retry_reason
+
+    @pytest.mark.backend("mysql", "postgres")
+    def test_concurrent_loop_gate_completions_create_one_successor(self, 
session, running_loop_gate):
+        gate = running_loop_gate
+        dr = gate.dag_run
+        gate_id = gate.id
+        bind = session.get_bind()
+        barrier = Barrier(2)
+
+        def complete():
+            with Session(bind=bind) as request_session:
+                barrier.wait(timeout=10)
+                result = ti_update_state(
+                    task_instance_id=gate_id,
+                    ti_patch_payload=TISuccessStatePayload(state="success", 
end_date=DEFAULT_END_DATE),
+                    session=request_session,
+                    dag_bag=DBDagBag(),
+                )
+                return result.status_code if result is not None else 204
+
+        with ThreadPoolExecutor(max_workers=2) as pool:
+            futures = [pool.submit(complete) for _ in range(2)]
+            assert sorted(future.result(timeout=20) for future in futures) == 
[200, 204]
+        session.expire_all()
+        assert gate.state == State.SUCCESS
+        assert sorted(
+            ti.region_index for ti in dr.get_task_instances(session=session) 
if ti.task_id == gate.task_id
+        ) == [0, 1]
+        assert session.scalar(select(XComModel).where(XComModel.task_id == 
gate.task_id)) is None
+
+    @pytest.mark.backend("mysql", "postgres")
+    def test_loop_decision_rewrite_and_completion_use_same_lock_order(
+        self, session, running_loop_gate, mocker
+    ):
+        gate = running_loop_gate
+        gate_id, dag_id, run_id, task_id = gate.id, gate.dag_id, gate.run_id, 
gate.task_id
+        bind = session.get_bind()
+        run_locked = Event()
+        completion_at_lock = Event()
+        execute = Session.execute
+
+        def coordinate_requests(request_session, statement, *args, **kwargs):
+            role = request_session.info.get("loop_test_role")
+            locks_run = (
+                isinstance(statement, Select)
+                and statement._for_update_arg is not None
+                and any(getattr(table, "name", None) == "dag_run" for table in 
statement.get_final_froms())
+            )
+            if role == "completion" and locks_run:
+                completion_at_lock.set()
+            result = execute(request_session, statement, *args, **kwargs)
+            if role == "writer" and locks_run and not run_locked.is_set():
+                run_locked.set()
+                assert completion_at_lock.wait(timeout=10)
+            return result
+
+        mocker.patch.object(Session, "execute", autospec=True, 
side_effect=coordinate_requests)
+
+        def rewrite():
+            with Session(bind=bind, info={"loop_test_role": "writer"}) as 
request_session:
+                set_xcom(
+                    dag_id=dag_id,
+                    run_id=run_id,
+                    task_id=task_id,
+                    key="_airflow_loop_decision",
+                    session=request_session,
+                    dag_bag=DBDagBag(),
+                    value="continue",

Review Comment:
   The writer only pauses after its `dag_run` row lock, and completion starts 
after that, so the writer already holds both locks before completion asks for 
one. If `set_xcom` locked the gate TI first, this would still pass. Pausing 
after the writer's first locking select, whichever table it is, would make a 
reversed order deadlock here. And since the fixture already seeds `"continue"`, 
writing `"continue"` again can't show whose decision completion consumed; 
seeding `"stop"` would.



##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -2177,6 +2186,121 @@ def create_ti(task: Operator, indexes: Iterable[int], 
region_id: UUID) -> Iterat
             creator = create_ti
         return creator
 
+    def complete_loop_gate(
+        self,
+        gate: TI,
+        group: SerializedLoopTaskGroup,
+        state: TaskInstanceState,
+        *,
+        session: Session,
+    ) -> None:
+        """Consume a gate decision while the caller holds the DagRun and TI 
locks."""
+        from airflow.models.xcom import XComModelV2
+        from airflow.settings import task_instance_mutation_hook
+
+        if gate.dag_version_id is None:
+            raise InvalidLoopDecision("Loop gate requires a pinned DAG 
version")
+        signal = XComModelV2.get_for_attempt(gate.id, LOOP_DECISION_KEY, 
session=session)
+        later_gate = session.scalar(
+            select(TI.id)
+            .where(
+                TI.working_set.is_(True),
+                TI.dag_id == self.dag_id,
+                TI.run_id == self.run_id,
+                TI.task_id == gate.task_id,
+                TI.region_id == gate.region_id,
+                TI.region_index > gate.region_index,
+            )
+            .limit(1)
+        )
+        if state != TaskInstanceState.SUCCESS or later_gate:
+            if signal is not None:
+                session.delete(signal)
+            return
+        decision = signal.value if signal is not None else None
+        at_limit = gate.region_index + 1 >= group.max_iterations
+        if decision not in ("continue", "stop"):
+            raise InvalidLoopDecision("Successful loop gate requires a 
decision")
+        if decision == "continue" and at_limit:
+            raise InvalidLoopDecision("Loop cannot continue beyond its 
iteration limit")
+        if not group.has_until and (decision == "stop") != at_limit:
+            raise InvalidLoopDecision("Fixed-count loop decision does not 
match its iteration limit")
+        session.delete(signal)
+        if decision == "stop":
+            return
+        created_counts: dict[str, int] = defaultdict(int)
+        hook_is_noop: Literal[True, False] = 
getattr(task_instance_mutation_hook, "is_noop", False)
+        creator = self._get_task_creator(
+            created_counts, task_instance_mutation_hook, hook_is_noop, 
gate.dag_version_id
+        )
+        tasks = list(
+            self._create_tasks(
+                group.iter_tasks(),
+                creator,
+                session=session,
+                parent_region=(gate.region_id, gate.region_index + 1),
+            )
+        )
+        self._create_task_instances(
+            self.dag_id, tasks, created_counts, hook_is_noop, session=session, 
propagate_errors=True
+        )
+
+    def _create_initial_tasks(
+        self,
+        tasks: Iterable[Operator],
+        task_creator: Callable[[Operator, Iterable[int], UUID], CreatedTasks],
+        *,
+        session: Session,
+    ) -> CreatedTasks:
+        from airflow.models.task_coordinates import enclosing_loop
+
+        ordinary_tasks = []
+        loop_tasks = defaultdict(list)
+        loops = {}
+        for task in tasks:
+            if loop := enclosing_loop(task):
+                loops[loop.group_id] = loop
+                loop_tasks[loop.group_id].append(task)
+            else:
+                ordinary_tasks.append(task)
+        yield from self._create_tasks(ordinary_tasks, task_creator, 
session=session, expand_literals=True)
+        for group_id, members in loop_tasks.items():
+            coordinates = (
+                session.execute(
+                    select(TI.region_id, TI.region_index).where(
+                        TI.working_set.is_(True),
+                        TI.dag_id == self.dag_id,
+                        TI.run_id == self.run_id,
+                        TI.task_id == loops[group_id].gate_task_id,
+                    )
+                )
+                .tuples()
+                .all()
+            )
+            if not coordinates:
+                if session.scalar(
+                    select(DynamicRegion.id)
+                    .where(
+                        DynamicRegion.dag_id == self.dag_id,
+                        DynamicRegion.run_id == self.run_id,
+                        DynamicRegion.node_id == group_id,
+                    )
+                    .limit(1)
+                ):
+                    raise ValueError(f"Loop {group_id!r} has regions but no 
live gate")
+                region = DynamicRegion(dag_id=self.dag_id, run_id=self.run_id, 
node_id=group_id)

Review Comment:
   `DynamicRegion.get_or_create` exists so the loser of a concurrent insert 
reuses the winner's row. Here two `verify_integrity` calls racing on the same 
run both pass the check above, and the second gets an `IntegrityError` on 
`slot_key` from the flush, before `_create_task_instances` gets a chance to 
handle it. Could this use `get_or_create`?



##########
task-sdk/tests/task_sdk/execution_time/test_loop.py:
##########
@@ -0,0 +1,225 @@
+# 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.
+from __future__ import annotations
+
+import pytest
+
+from airflow.sdk import DAG, BaseOperator, TaskInstanceState, task, task_group
+from airflow.sdk.api.datamodels._generated import LoopContext
+from airflow.sdk.bases.xcom import BaseXCom
+from airflow.sdk.definitions._internal.loop import create_loop
+from airflow.sdk.exceptions import AirflowFailException
+from airflow.sdk.execution_time import task_runner
+from airflow.sdk.execution_time.comms import (
+    GetXCom,
+    GetXComCount,
+    GetXComSequenceItem,
+    GetXComSequenceSlice,
+    SetXCom,
+    XComCountResponse,
+    XComResult,
+    XComSequenceIndexResult,
+    XComSequenceSliceResult,
+)
+from airflow.sdk.execution_time.lazy_sequence import LazyXComSequence
+
+
[email protected]
+def loop_ti(create_runtime_ti):
+    def make(*, index=0, max_iterations=3, until=None, map_index=-1):
+        @task_group
+        def body():
+            BaseOperator(task_id="terminal")
+
+        with DAG("loop_runtime", schedule=None) as dag:
+            group = create_loop(body, max_iterations=max_iterations, 
until=until)
+        ti = create_runtime_ti(task=dag.get_task(group.gate_task_id), 
map_index=map_index)
+        ti._ti_context_from_server.loop = LoopContext(
+            node_id=group.group_id,
+            index=index,
+            max_iterations=max_iterations,
+            terminal_task_id=group.terminal_task_id,
+            terminal_is_mapped=False,
+        )
+        return ti
+
+    return make
+
+
+def test_first_iteration_previous_does_not_read_xcom(loop_ti, 
mock_supervisor_comms):
+    ti = loop_ti(map_index=7)
+    loop = ti.get_template_context()["loop"]
+
+    assert loop.index == 0
+    assert loop.max_iterations == 3
+    assert loop.previous is None
+    assert ti.map_index == 7
+    mock_supervisor_comms.send.assert_not_called()
+
+
[email protected]("mapped", [False, True])
+def 
test_body_callable_receives_loop_context_after_unmapping(create_runtime_ti, 
mapped):
+    @task
+    def terminal(value, *, loop, ti):
+        return value, loop.index, ti.map_index
+
+    @task_group
+    def body():
+        if mapped:
+            terminal.expand(value=[5])
+        else:
+            terminal(5)
+
+    with DAG("loop_body_runtime", schedule=None) as dag:
+        group = create_loop(body, max_iterations=4)
+    operator = dag.get_task(group.terminal_task_id)
+    if mapped:
+        operator = operator.unmap({"op_kwargs": {"value": 5}})
+    ti = create_runtime_ti(task=operator, map_index=0 if mapped else -1)
+    ti._ti_context_from_server.loop = LoopContext(
+        node_id=group.group_id,
+        index=2,
+        max_iterations=4,
+        terminal_task_id=group.terminal_task_id,
+        terminal_is_mapped=mapped,
+    )
+
+    assert operator.execute(ti.get_template_context()) == (5, 2, 0 if mapped 
else -1)
+
+
[email protected]("value", [0, False, [], ""])
+def test_previous_and_current_result_preserve_falsey_values(loop_ti, 
mock_supervisor_comms, value):
+    ti = loop_ti(index=1)
+    mock_supervisor_comms.send.return_value = XComResult(key="return_value", 
value=value)
+    loop = ti.get_template_context()["loop"]
+
+    assert loop.previous == value
+    previous = mock_supervisor_comms.send.call_args.args[0]
+    assert isinstance(previous, GetXCom)
+    assert previous.task_id == "body.terminal"
+    assert previous.previous_iteration is True
+    assert loop.result == value
+    assert mock_supervisor_comms.send.call_args.args[0].previous_iteration is 
False
+
+
[email protected](
+    ("index", "condition", "decision"),
+    [(0, None, "continue"), (2, None, "stop"), (0, False, "continue"), (0, 
True, "stop"), (2, True, "stop")],
+)
+def test_gate_publishes_successful_decision(loop_ti, mock_supervisor_comms, 
index, condition, decision):
+    def until(*, loop):
+        assert loop.index == index
+        return condition
+
+    ti = loop_ti(index=index, until=until if condition is not None else None)
+    context = ti.get_template_context()
+
+    ti.task.execute(context)
+
+    message = mock_supervisor_comms.send.call_args.args[0]
+    assert isinstance(message, SetXCom)
+    assert message.key == "_airflow_loop_decision"
+    assert message.value == decision
+    assert message.loop_decision is True
+
+
[email protected]("raises", [False, True])
+def test_unsuccessful_gate_does_not_publish_decision(loop_ti, 
mock_supervisor_comms, raises):
+    def until(*, loop):
+        if raises:
+            raise ValueError("condition failed")
+        return False
+
+    ti = loop_ti(index=2, until=until)
+    with pytest.raises(Exception, match="condition failed" if raises else 
"max_iterations"):
+        ti.task.execute(ti.get_template_context())
+
+    mock_supervisor_comms.send.assert_not_called()
+
+
[email protected]("result", [None, 0, 1, [], "yes"])
+def test_gate_fails_when_until_does_not_return_a_bool(loop_ti, 
mock_supervisor_comms, result):
+    ti = loop_ti(index=0, until=lambda: result)
+
+    with pytest.raises(AirflowFailException, match=f"got 
{type(result).__name__}"):
+        ti.task.execute(ti.get_template_context())
+
+    mock_supervisor_comms.send.assert_not_called()
+
+
[email protected]("previous_iteration", [False, True])
[email protected]("values", [[], [0], [0, False, ""]])
+def test_mapped_terminal_keeps_iteration_across_lazy_reads(
+    loop_ti, mock_supervisor_comms, previous_iteration, values
+):
+    ti = loop_ti(index=1)
+    ti._ti_context_from_server.loop.terminal_is_mapped = True
+
+    def respond(message):
+        assert message.previous_iteration is previous_iteration
+        if isinstance(message, GetXComCount):
+            return XComCountResponse(len=len(values))
+        if isinstance(message, GetXComSequenceItem):
+            return XComSequenceIndexResult(root=values[message.offset])
+        assert isinstance(message, GetXComSequenceSlice)
+        return XComSequenceSliceResult(root=values[slice(message.start, 
message.stop, message.step)])
+
+    mock_supervisor_comms.send.side_effect = respond
+    context = ti.get_template_context()["loop"]
+    result = context.previous if previous_iteration else context.result
+
+    assert isinstance(result, LazyXComSequence)
+    assert len(result) == len(values)
+    assert result[:] == values
+    assert result[::-1] == values[::-1]
+    if values:
+        assert result[-1] == values[-1]
+
+
+def test_loop_decision_bypasses_custom_backend(loop_ti, mock_supervisor_comms, 
mocker):
+    serialize = 
mocker.patch("airflow.sdk.execution_time.xcom.XCom.serialize_value", 
autospec=True)
+    ti = loop_ti()
+
+    ti.task.execute(ti.get_template_context())
+
+    serialize.assert_not_called()
+    assert mock_supervisor_comms.send.call_args.args[0].value == "continue"
+
+
+def test_gate_retry_clears_old_signal_without_custom_backend(loop_ti, 
mock_supervisor_comms, monkeypatch):
+    class CustomXCom(BaseXCom):
+        @classmethod
+        def purge(cls, xcom, *args):
+            raise AssertionError("Control metadata must not reach custom 
backend")
+
+    ti = loop_ti()
+    ti._ti_context_from_server.xcom_keys_to_clear = ["_airflow_loop_decision"]
+    monkeypatch.setattr(task_runner, "XCom", CustomXCom)
+    mock_supervisor_comms.send.side_effect = lambda message: (
+        XComResult(key="_airflow_loop_decision", value="stop") if 
isinstance(message, GetXCom) else None
+    )
+
+    state, _, error = task_runner.run(ti, context=ti.get_template_context(), 
log=ti.task.log)
+
+    assert error is None
+    assert state == TaskInstanceState.SUCCESS
+    decisions = [
+        call.args[0]
+        for call in mock_supervisor_comms.send.call_args_list
+        if call.args and isinstance(call.args[0], SetXCom)
+    ]
+    assert [decision.value for decision in decisions] == ["continue"]

Review Comment:
   This shows the custom backend isn't used, but nothing checks the old 
decision is deleted; if the loop-key branch skipped the delete entirely, this 
would still pass. Asserting a `DeleteXCom` for `_airflow_loop_decision` in 
`send.call_args_list` would cover it.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to