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


##########
airflow-core/src/airflow/models/task_coordinates.py:
##########
@@ -221,38 +241,110 @@ def has_regions(self, dag_id: str, run_id: str | None, 
task_id: str) -> bool:
             query = query.where(TaskInstance.run_id == run_id)
         return bool(self.session.scalar(select(query.exists())))
 
-    def resolve(
+    def resolve_dependency(self, caller: TaskInstance, task_id: str) -> 
tuple[TaskInstance, ...]:
+        producer = self.get_task(caller.dag_id, caller.run_id, task_id, 
dag_version_id=caller.dag_version_id)

Review Comment:
   This producer `get_task` isn't guarded, and `upstream_tis` now always comes 
through here. Before, `resolve()` either took the legacy lookup or caught 
`TaskNotFound` in `_build_producer_request` and fell back to the removed-task 
query.
   
   One way to hit it: a Dag with any mapped task (so the run has regions) on an 
unversioned bundle, and a downstream task T that is still in the NULL state. 
The scheduler's re-pin only updates `State.unfinished`, and `IN (NULL, ...)` 
never matches NULL, so T stays on v1. The user then adds a new upstream X to T. 
`ti.task` comes from the latest Dag, so the trigger rule dep asks for X, and 
`get_task("X", v1)` raises. `_schedule_all_dag_runs` logs "Error scheduling DAG 
run" and the run stops advancing on every loop until someone clears T.
   
   Could this catch `TaskNotFound` and fall through to `self.resolve(...)`, the 
way `select_skip_target_ids` already does? Only the gate branch needs the 
producer definition.



##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -430,6 +444,32 @@ def ti_update_state(
             raise HTTPException(status_code=409, detail={"reason": 
"invalid_state"})
         return Response(status_code=status.HTTP_204_NO_CONTENT)
 
+    loop_group = None
+    loop_gate = None
+    if isinstance(ti_patch_payload, (TISuccessStatePayload, 
TITerminalStatePayload)):
+        gate_run = session.execute(
+            select(TI.dag_id, TI.run_id).where(
+                TI.id == task_instance_id,
+                TI.working_set.is_(True),
+                TI.operator == "LoopGateOperator",
+            )
+        ).one_or_none()
+        if gate_run is not None:
+            session.execute(
+                select(DR).where(DR.dag_id == gate_run.dag_id, DR.run_id == 
gate_run.run_id).with_for_update()
+            ).scalar_one()
+            loop_gate = session.scalar(
+                select(TI)
+                .where(TI.id == task_instance_id, TI.working_set.is_(True))
+                .with_for_update(of=TI)
+                .execution_options(populate_existing=True)
+            )
+            if loop_gate is not None:
+                loop_context = TaskCoordinateResolver(dag_bag, 
session).loop_context(loop_gate)
+                if loop_context is None or loop_context[0].gate_task_id != 
loop_gate.task_id:

Review Comment:
   Gates are picked by `TI.operator == "LoopGateOperator"`, which is just the 
class name, so a user operator with that name outside a loop lands in this 
branch. `loop_context` returns None for the sentinel region, and every success 
or failure PATCH gets a 409. The supervisor discards a 409'd outcome, so the 
task never succeeds and ends up failed through the executor state-mismatch 
path, on every retry.
   
   Could `loop_context is None` fall through to the normal path instead, with 
the 409 kept for a task that is inside a loop but isn't its gate? The same 
string is checked in the scheduler (`scheduler_job_runner.py:1682`), so a 
shared constant or helper would keep the two in step.



##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -517,6 +557,25 @@ def ti_update_state(
                 detail={"reason": "invalid_partition_key", "message": str(e)},
             ) from e
 
+    gate_completed = False
+    if (
+        loop_gate is not None
+        and loop_group is not None
+        and isinstance(ti_patch_payload, TISuccessStatePayload)
+    ):
+        try:
+            loop_gate.dag_run.complete_loop_gate(
+                loop_gate, loop_group, TaskInstanceState.SUCCESS, 
session=session
+            )
+            gate_completed = True
+        except InvalidLoopDecision as error:
+            log.warning("Loop gate success rejected", error=str(error))
+            ti_patch_payload = _build_rejected_gate_payload(

Review Comment:
   When this rewrites the payload to retry or failed, the worker still gets a 
204 and has already gone on to `finalize(state=SUCCESS)`. So 
`on_success_callback` and the success listeners run, and `on_failure_callback`, 
`on_retry_callback` and failure emails never do, while the gate log ends with 
"Loop iteration N: continue".
   
   A reachable case: an unversioned bundle where the author lowers 
`max_iterations` while a gate is running. The scheduler re-pins RUNNING TIs to 
the new version, so `complete_loop_gate` checks against the new limit and 
rejects the decision the worker computed from the old one.
   
   Could the `at_limit` and fixed-count checks also run when the decision is 
written in `set_xcom` (it already loads the pinned loop there)? Then the SDK's 
`SetXCom` fails inside `execute()`, the normal retry and failure handling runs 
the right callbacks, and this path stays as a backstop.



##########
airflow-core/src/airflow/jobs/scheduler_job_runner.py:
##########
@@ -1679,6 +1679,10 @@ def process_executor_events(
 
                     task = dag.get_task(ti.task_id)
                 except Exception:
+                    if state == TaskInstanceState.SUCCESS and ti.operator == 
"LoopGateOperator":
+                        ti.task = None
+                        ti.handle_failure(error=msg, session=session)

Review Comment:
   This `continue`s before the `cls.logger().exception(...)` below, so when the 
gate's task can't be loaded (a `TaskNotFound` after the gate id changes in a 
new Dag version, say) nothing records why. `msg` only describes the state 
mismatch. Could this log the exception before `handle_failure`, or append it to 
`msg`?



##########
task-sdk/src/airflow/sdk/execution_time/loop.py:
##########
@@ -0,0 +1,103 @@
+# 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
+
+from typing import TYPE_CHECKING, Any
+
+import attrs
+
+from airflow.sdk.bases.decorator import determine_kwargs
+from airflow.sdk.definitions._internal.loop import LOOP_DECISION_KEY
+from airflow.sdk.definitions.xcom_arg import PlainXComArg
+from airflow.sdk.exceptions import AirflowException
+from airflow.sdk.execution_time.comms import SetXCom
+from airflow.sdk.execution_time.lazy_sequence import LazyXComSequence
+from airflow.sdk.execution_time.xcom import XCom
+
+if TYPE_CHECKING:
+    from airflow.sdk.api.datamodels._generated import LoopContext
+    from airflow.sdk.definitions._internal.loop import LoopGateOperator
+    from airflow.sdk.definitions.context import Context
+    from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance
+
+
+class LoopMaxIterationsExceeded(AirflowException):

Review Comment:
   Since this is an `AirflowException`, the cap failure is retryable, and the 
gate inherits `default_args` retries like any other operator. With `retries=3` 
and the default `retry_delay`, a loop that never converges sits in up_for_retry 
for about 15 minutes and calls `until` three more times against the same 
`loop.result` before it reaches the failed state loops.rst describes.
   
   Would subclassing `AirflowFailException` fit better? Exceptions raised by 
`until` itself would stay retryable. If keeping retries is deliberate (an 
`until` that polls external state), the docs should say the cap failure follows 
the gate's retries.



##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -2132,6 +2141,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:
   The coordinates here come from TIs of the gate id in the new Dag version. On 
an unversioned bundle, renaming the `until` function mid-run (the gate id is 
`until.__name__`), or adding `until=` to a fixed-count loop, changes that id. 
This query then finds no gate, the loop already has regions, and we raise.
   
   By then `_verify_integrity_if_dag_changed` has re-pinned the unfinished TIs 
to the new version, and `_schedule_all_dag_runs` swallows the exception and 
still commits. So `check_version_id_exists_in_dr` is true on every later loop 
and `verify_integrity` never runs again. The new gate is never created, the old 
live gate points at a version without its task (`loop_context` raises 
`TaskNotFound` in `ti_update_state`), and the only trace is one "Error 
scheduling DAG run" log line.
   
   Could the live passes come from the loop's regions or member TIs 
(`DynamicRegion.node_id == group_id`) instead of gate TIs, or at least skip 
just this loop with a warning rather than raising inside the scheduler's 
catch-all?
   
   Related question: a newly added member is created at every live gate 
coordinate, so a task added while the loop is at pass 3 also runs in passes 0, 
1 and 2, whose gates already succeeded. Is that intended? 
`test_loop_new_member_joins_live_pass_after_reserialization` only has one pass, 
so it can't tell the two behaviours apart.



##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -540,6 +599,8 @@ def ti_update_state(
         # Let DataErrorHandler return a 422 instead of silently marking the TI 
FAILED below.
         raise
     except Exception:
+        if loop_gate is not None:
+            raise

Review Comment:
   This re-raise changes what a gate gets on an unexpected error here (a 500 
with the gate left RUNNING, instead of rollback plus FAILED), and no test 
reaches it. Both atomicity tests inject the failure inside 
`complete_loop_gate`, so nothing fails after the next pass has been flushed and 
before the gate UPDATE runs. A `session.commit()` slipped in right after 
`complete_loop_gate` would keep every current test green.
   
   Could one test patch `_create_ti_state_update_query_and_update_state` 
(autospec, raising) on `running_loop_gate` and assert the 500, the gate still 
RUNNING, no pass-1 rows, and the decision row still there?



##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py:
##########
@@ -289,6 +304,7 @@ class GetXComSliceFilterParams(BaseModel):
     include_prior_dates: bool = False
     region_id: UUID | None = None
     region_index: int | None = None
+    previous_iteration: bool = False

Review Comment:
   `previous_iteration` is added to `GetXComSliceFilterParams` and 
`GetXcomFilterParams` without a Cadwyn instruction, so the generated spec for 
the older API versions now lists it too. The last field added to these models 
got one (`AddIncludePriorDatesToGetXComSlice` in `v2025_08_10.py`). Could 
`AddLoopContext`, or a small new change in `v2026_10_30.py`, add 
`schema(...).field("previous_iteration").didnt_exist` for both? The prek 
version check only watches `datamodels/`, which is why nothing flagged it.



##########
airflow-core/src/airflow/models/dynamic_region.py:
##########
@@ -216,62 +190,220 @@ def resolve_current_producers(
             for ref in (row.parent_region_id, row.forked_from_region_id)
             if ref is not None and ref not in regions
         }
+    return regions
+
+
+def loop_position(
+    regions: dict[UUID, DynamicRegion], coordinate_id: UUID, index: int, 
loop_node_id: str
+) -> tuple[UUID, int] | None:
+    """Return the enclosing loop's fork family and iteration."""
+    seen: set[UUID] = set()
+    while coordinate_id != SENTINEL_REGION_ID:
+        if coordinate_id in seen:
+            raise ValueError("Cyclic region ancestry")
+        seen.add(coordinate_id)
+        region = regions[coordinate_id]
+        if region.node_id == loop_node_id:
+            family = region
+            lineage: set[UUID] = set()
+            while family.forked_from_region_id is not None:
+                if family.id in lineage:
+                    raise ValueError("Cyclic region lineage")
+                lineage.add(family.id)
+                family = regions[family.forked_from_region_id]
+            return family.id, index
+        if region.parent_region_id is None:
+            break
+        if TYPE_CHECKING:
+            assert region.parent_region_index is not None
+        coordinate_id, index = region.parent_region_id, 
region.parent_region_index
+    return None
+
 
-    loop_node_id = context.loop_node_id if context else None
-
-    def loop_position(coordinate_id: UUID, index: int) -> tuple[UUID, int] | 
None:
-        seen: set[UUID] = set()
-        while coordinate_id != zero:
-            if coordinate_id in seen:
-                raise ValueError("Cyclic region ancestry")
-            seen.add(coordinate_id)
-            region = regions[coordinate_id]
-            if region.node_id == loop_node_id:
-                family = region
-                lineage: set[UUID] = set()
-                while family.forked_from_region_id is not None:
-                    if family.id in lineage:
-                        raise ValueError("Cyclic region lineage")
-                    lineage.add(family.id)
-                    family = regions[family.forked_from_region_id]
-                return family.id, index
-            if region.parent_region_id is None:
-                break
-            if TYPE_CHECKING:
-                assert region.parent_region_index is not None
-            coordinate_id, index = region.parent_region_id, 
region.parent_region_index
-        return None
+def _filter_producers(
+    query: Select,
+    *,
+    dag_id: str,
+    run_id: str,
+    task_id: str,
+    is_mapped: bool,
+    map_indexes: int | Collection[int] | None,
+    region_id: UUID | None,
+    region_index: int | None,
+    top_level_only: bool,
+) -> Select:
+    """Push every coordinate predicate that does not need loop ancestry into 
SQL."""
+    from airflow.models.taskinstance import TaskInstance
+
+    query = query.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 top_level_only:
+        if is_mapped:
+            top_level_region = (
+                select(DynamicRegion.id)
+                .where(DynamicRegion.id == TaskInstance.region_id, 
DynamicRegion.parent_region_id.is_(None))
+                .correlate(TaskInstance)
+                .exists()
+            )
+            query = query.where(or_(TaskInstance.region_id == 
SENTINEL_REGION_ID, top_level_region))
+        else:
+            query = query.where(TaskInstance.region_id == SENTINEL_REGION_ID, 
TaskInstance.region_index == -1)
+    if map_indexes is None:
+        return query
+    if not is_mapped:
+        wanted = map_indexes == -1 if isinstance(map_indexes, int) else -1 in 
map_indexes
+        return query if wanted else query.where(false())
+    if isinstance(map_indexes, int):
+        return query.where(TaskInstance.region_index == map_indexes)
+    if isinstance(map_indexes, range) and map_indexes.step == 1:
+        return query.where(
+            TaskInstance.region_index >= map_indexes.start, 
TaskInstance.region_index < map_indexes.stop
+        )
+    return query.where(TaskInstance.region_index.in_(list(map_indexes)))
 
+
+def _validate_producer_request(
+    *, context: ProducerContext | None, region_id: UUID | None, region_index: 
int | None
+) -> None:
+    if region_index is not None and region_id is None:
+        raise ValueError("region_index requires an explicit producer 
region_id")
+    if context and context.previous_iteration and context.loop_node_id is None:
+        raise ValueError("Previous-iteration lookup requires a loop context")
+
+
+def resolve_current_producers(
+    *,
+    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,
+) -> tuple[TaskInstance, ...]:
+    """Resolve the live producer task instances whose data the caller reads by 
task instance UUID."""
+    from airflow.models.taskinstance import TaskInstance
+
+    _validate_producer_request(context=context, region_id=region_id, 
region_index=region_index)
+    regions: dict[UUID, DynamicRegion] = {}
     position = None
-    if context and context.loop_node_id is not None and region_id is None:
-        position = loop_position(context.region_id, context.region_index)
-        if position is None:
-            raise ValueError("Caller is not inside the requested loop")
-        if context.previous_iteration:
-            position = position[0], position[1] - 1
-            if position[1] < 0:
-                return ()
+    if context and region_id is None:
+        regions = load_region_ancestry(
+            {context.region_id} - {SENTINEL_REGION_ID}, dag_id=dag_id, 
run_id=run_id, session=session
+        )
+        if context.loop_node_id is not None:
+            position = loop_position(regions, context.region_id, 
context.region_index, context.loop_node_id)
+            if position is None:
+                raise ValueError("Caller is not inside the requested loop")
+            if context.previous_iteration:
+                position = position[0], position[1] - 1
+                if position[1] < 0:
+                    return ()
+
+    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,
+        )
+    ).all()
+    if position is not None:
+        if TYPE_CHECKING:
+            assert context is not None
+            assert context.loop_node_id is not None
+        regions.update(
+            load_region_ancestry(
+                {ti.region_id for ti in candidates} - set(regions) - 
{SENTINEL_REGION_ID},
+                dag_id=dag_id,
+                run_id=run_id,
+                session=session,
+            )
+        )
+        candidates = [
+            ti
+            for ti in candidates
+            if loop_position(regions, ti.region_id, ti.region_index, 
context.loop_node_id) == position
+        ]
 
     selected: dict[int, TaskInstance] = {}
     for ti in candidates:
-        if region_id is None:
-            if position is not None:
-                if loop_position(ti.region_id, ti.region_index) != position:
-                    continue
-            elif ti.region_id != zero:
-                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

Review Comment:
   With a context set, this always falls back to `resolve_current_producers`, 
which loads every live TI of the producer across all passes as full ORM rows 
and narrows to the right pass in Python. For `loop.result` on a mapped terminal 
that adds up: `LazyXComSequence` fetches one item per request, and each 
`/item/{offset}` read comes through here with no `map_index`.
   
   The documented `min(loop.result)` over M mapped slots at pass K is M+1 
requests that each hydrate K*M TaskInstance rows (plus the joined DagRun). At 
M=1000 and K=10 that's about 10M row loads for one gate run, and it grows every 
pass because earlier passes stay live.
   
   Could the pass filter go into SQL first (the loop region at `region_index == 
k`, plus child regions with `parent_region_index == k`), selecting only `id, 
region_id, region_index`, with the Python `loop_position` check kept as the 
exact filter over that smaller set?



##########
airflow-core/src/airflow/models/task_coordinates.py:
##########
@@ -221,38 +241,110 @@ def has_regions(self, dag_id: str, run_id: str | None, 
task_id: str) -> bool:
             query = query.where(TaskInstance.run_id == run_id)
         return bool(self.session.scalar(select(query.exists())))
 
-    def resolve(
+    def resolve_dependency(self, caller: TaskInstance, task_id: str) -> 
tuple[TaskInstance, ...]:
+        producer = self.get_task(caller.dag_id, caller.run_id, task_id, 
dag_version_id=caller.dag_version_id)
+        loop = enclosing_loop(producer)
+        caller_task = self.get_task(
+            caller.dag_id, caller.run_id, caller.task_id, 
dag_version_id=caller.dag_version_id
+        )
+        caller_loop = enclosing_loop(caller_task)
+        if (
+            loop is not None
+            and loop.gate_task_id == task_id
+            and (caller_loop is None or caller_loop.group_id != loop.group_id)
+        ):
+            gates = self.session.scalars(
+                select(TaskInstance)
+                .join(DynamicRegion, DynamicRegion.id == 
TaskInstance.region_id)
+                .where(
+                    TaskInstance.working_set.is_(True),
+                    TaskInstance.dag_id == caller.dag_id,
+                    TaskInstance.run_id == caller.run_id,
+                    TaskInstance.task_id == task_id,
+                    DynamicRegion.node_id == loop.group_id,
+                )
+                .order_by(TaskInstance.region_index.desc())
+                .limit(2)
+            ).all()
+            if len(gates) == 2 and gates[0].region_index == 
gates[1].region_index:
+                raise AmbiguousProducerError(f"Multiple live loop gates at 
pass {gates[0].region_index}")
+            return tuple(gates[:1])
+        return self.resolve(dag_id=caller.dag_id, run_id=caller.run_id, 
task_id=task_id, caller=caller)
+
+    def _filter_legacy_producers(
         self,
+        query: Select,
         *,
         dag_id: str,
         run_id: str,
         task_id: str,
-        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, ...]:
-        if region_index is not None and region_id is None:
-            raise ValueError("region_index requires an explicit producer 
region_id")
-        if not previous_iteration and (
-            region_id == UUID(int=0) or (region_id is None and not 
self.has_regions(dag_id, run_id, task_id))
-        ):
-            query = select(TaskInstance).where(
-                TaskInstance.working_set.is_(True),
-                TaskInstance.dag_id == dag_id,
-                TaskInstance.run_id == run_id,
-                TaskInstance.task_id == task_id,
-                TaskInstance.region_id == UUID(int=0),
-            )
-            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)))
+        region_index: int | None,
+        map_indexes: int | Collection[int] | None,
+    ) -> Select:
+        query = query.where(
+            TaskInstance.working_set.is_(True),
+            TaskInstance.dag_id == dag_id,
+            TaskInstance.run_id == run_id,
+            TaskInstance.task_id == task_id,
+            TaskInstance.region_id == SENTINEL_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 query
 
+    def _filter_removed_task_producers(
+        self,
+        query: Select,
+        *,
+        dag_id: str,
+        run_id: str,
+        task_id: str,
+        region_id: UUID | None,
+        region_index: int | None,
+        map_indexes: int | Collection[int] | None,
+    ) -> Select:
+        """Find live rows of a task whose definition is gone, using only its 
own expansion regions."""
+        query = query.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_id is not None:
+            query = query.where(TaskInstance.region_id == region_id)
+        if region_index is not None:

Review Comment:
   `_filter_legacy_producers` and `_filter_removed_task_producers` repeat the 
same `region_index` / `map_indexes` clause block line for line, and 
`_build_producer_request` returns a `dict[str, Any]` that is splatted into two 
different functions, so mypy no longer checks those kwargs. A small shared 
clause helper and a NamedTuple for the request would keep the copies from 
drifting.



##########
airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py:
##########
@@ -1460,6 +1516,457 @@ 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
+
+    def test_loop_decision_rewrite_and_completion_use_same_lock_order(

Review Comment:
   This one has no `@pytest.mark.backend("mysql", "postgres")`, unlike the 
concurrency test just above. On SQLite `with_for_update()` renders to nothing, 
so the writer never blocks the completion and the result depends on which 
thread writes first. If the completion commits before the writer's upsert, the 
writer has already passed its RUNNING check and re-creates the decision row, 
and the `is None` assertion fails. It also can't catch an inverted lock order 
on SQLite, since nothing deadlocks there.



##########
airflow-core/tests/unit/models/test_task_coordinates.py:
##########
@@ -198,3 +221,75 @@ def 
test_unversioned_run_keeps_task_definition_after_latest_graph_changes(
             region_id=mapped.region_id,
             region_index=0,
         ) == (mapped,)
+
+
[email protected]
+def mapped_run(dag_maker, session):
+    def create(mapped_count: int):
+        with dag_maker(serialized=True):
+            mapped = PythonOperator.partial(task_id="mapped", 
python_callable=str).expand(
+                op_args=[[i] for i in range(mapped_count)]
+            )
+            mapped >> EmptyOperator(task_id="reduce")
+        dr = dag_maker.create_dagrun()
+        caller = session.scalars(
+            select(TaskInstance).where(TaskInstance.run_id == dr.run_id, 
TaskInstance.task_id == "reduce")
+        ).one()
+        session.expire_all()
+        return dr, caller
+
+    return create
+
+
+@contextmanager
+def count_loaded_task_instances(task_id: str):
+    loaded: list[TaskInstance] = []
+    listener = loaded.append
+    event.listen(TaskInstance, "load", listener)
+    try:
+        yield loaded
+    finally:
+        event.remove(TaskInstance, "load", listener)
+    loaded[:] = [ti for ti in loaded if ti.task_id == task_id]
+
+
[email protected]("mapped_count", [3, 60])
+def test_xcom_read_of_one_mapped_slot_loads_no_producer_rows(mapped_run, 
session, mapped_count):
+    dr, caller = mapped_run(mapped_count)
+    token = TIToken(id=caller.id, claims=TIClaims())
+
+    with count_loaded_task_instances("mapped") as loaded, 
assert_queries_count(8):
+        read = _build_xcom_read(
+            dag_id=dr.dag_id,
+            run_id=dr.run_id,
+            task_id="mapped",
+            key="return_value",
+            session=session,
+            dag_bag=DBDagBag(),
+            token=token,
+            map_index=2,
+        )
+        session.scalars(read).all()
+
+    assert loaded == []
+
+
+def test_selected_mapped_producers_match_resolved_producers(dag_maker, 
session):

Review Comment:
   This re-inlines the same Dag and caller lookup that the `mapped_run` fixture 
above builds, so `mapped_run(6)` would do. The `"load"` listener in 
`test_xcom_arg.py` also hand-writes what `count_loaded_task_instances` does 
here; that helper could move to `tests_common` and serve both.



-- 
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