This is an automated email from the ASF dual-hosted git repository.

ashb pushed a commit to branch task-loops-stack-3
in repository https://gitbox.apache.org/repos/asf/airflow.git

commit c65dc018088b488fa1a47b5478c33d4cfff5b2b1
Author: Ash Berlin-Taylor <[email protected]>
AuthorDate: Sun Oct 4 08:45:48 2026 +0100

    Resolve live producers by region and read their XCom by attempt
    
    Once a loop body or a mapped region can hold several live task instances 
with
    the same task_id and map_index, a consumer can no longer find its producer 
from
    (dag_id, run_id, task_id, map_index) alone, so the lookup has to know where 
the
    caller sits: in the same iteration, in the previous one (what a loop body 
means
    by "the last result"), outside the loop, or in an explicitly named region. 
When
    more than one live candidate still fits we have no sensible option to raise
    an error.
    
    Only live (non-archived or superceded) try are candidates, and their data is
    read by UUID. An archived try keeps its XCom under its own UUID, so reading 
by
    the resolved try is exact and cannot revive data from work a clear replaced.
    
    Callers that predate regions must see what they saw before, so the default 
read
    scope stays the sentinel region and regional rows appear only when a caller 
asks
    for them. Lookups of earlier runs refuse a non-sentinel region because a 
region
    belongs to a single Dag run, so the producer has to be resolved again for 
each
    run.
    
    The scheduler detected changes to upstream state and tracked map-length
    revisions by TaskInstanceKey, which cannot tell two regions apart. One 
region's
    expansion would have marked another's as changed, so both now key on the 
attempt
    and on (task, region).
---
 airflow-core/src/airflow/models/dagrun.py          |  53 ++-
 airflow-core/src/airflow/models/dynamic_region.py  | 139 +++++++-
 airflow-core/src/airflow/models/xcom.py            |  24 +-
 .../airflow/serialization/definitions/xcom_arg.py  | 101 ++++--
 .../tests/unit/models/test_dynamic_region.py       | 378 +++++++++++++++++++++
 airflow-core/tests/unit/models/test_xcom_arg.py    | 151 ++++++++
 6 files changed, 804 insertions(+), 42 deletions(-)

diff --git a/airflow-core/src/airflow/models/dagrun.py 
b/airflow-core/src/airflow/models/dagrun.py
index c929013f3a4..74f057f7765 100644
--- a/airflow-core/src/airflow/models/dagrun.py
+++ b/airflow-core/src/airflow/models/dagrun.py
@@ -124,7 +124,6 @@ if TYPE_CHECKING:
         TaskInstance as TIDataModel,
     )
     from airflow.models.dag_version import DagVersion
-    from airflow.models.taskinstancekey import TaskInstanceKey
     from airflow.sdk import DAG as SDKDAG
     from airflow.serialization.definitions.dag import SerializedDAG
     from airflow.serialization.definitions.mappedoperator import Operator
@@ -1064,6 +1063,7 @@ class DagRun(Base, LoggingMixin):
         task_id: str,
         *,
         map_index: int = -1,
+        region_id: UUID = SENTINEL_REGION_ID,
         session: Session = NEW_SESSION,
     ) -> TI | None:
         """
@@ -1078,6 +1078,7 @@ class DagRun(Base, LoggingMixin):
             task_id=task_id,
             session=session,
             map_index=map_index,
+            region_id=region_id,
         )
 
     @staticmethod
@@ -1088,6 +1089,7 @@ class DagRun(Base, LoggingMixin):
         task_id: str,
         *,
         map_index: int = -1,
+        region_id: UUID = SENTINEL_REGION_ID,
         session: Session = NEW_SESSION,
     ) -> TI | None:
         """
@@ -1099,7 +1101,9 @@ class DagRun(Base, LoggingMixin):
         :param session: Sqlalchemy ORM Session
         """
         return session.scalars(
-            select(TI).filter_by(dag_id=dag_id, run_id=dag_run_id, 
task_id=task_id, map_index=map_index)
+            select(TI).filter_by(
+                dag_id=dag_id, run_id=dag_run_id, task_id=task_id, 
map_index=map_index, region_id=region_id
+            )
         ).one_or_none()
 
     def get_dag(self) -> SerializedDAG:
@@ -1683,7 +1687,7 @@ class DagRun(Base, LoggingMixin):
         finished_tis: list[TI],
         session: Session,
     ) -> tuple[list[TI], bool, bool]:
-        old_states: dict[TaskInstanceKey, Any] = {}
+        old_states: dict[UUID, Any] = {}
         ready_tis: list[TI] = []
         changed_tis = False
 
@@ -1735,13 +1739,13 @@ class DagRun(Base, LoggingMixin):
         # Check dependencies.
         expansion_happened = False
         # Set of task ids for which was already done 
_revise_map_indexes_if_mapped
-        revised_map_index_task_ids: set[str] = set()
+        revised_map_index_task_ids: set[tuple[str, UUID]] = set()
         for schedulable in itertools.chain(schedulable_tis, additional_tis):
             if TYPE_CHECKING:
                 assert isinstance(schedulable.task, Operator)
             old_state = schedulable.state
             if not schedulable.are_dependencies_met(session=session, 
dep_context=dep_context):
-                old_states[schedulable.key] = old_state
+                old_states[schedulable.id] = old_state
                 continue
             # If schedulable is not yet expanded, try doing it now. This is
             # called in two places: First and ideally in the mini scheduler at
@@ -1762,12 +1766,16 @@ class DagRun(Base, LoggingMixin):
             if new_tis is None and schedulable.state in SCHEDULEABLE_STATES:
                 # It's enough to revise map index once per task id,
                 # checking the map index for each mapped task significantly 
slows down scheduling
-                if schedulable.task.task_id not in revised_map_index_task_ids:
+                expansion_key = (schedulable.task.task_id, 
schedulable.region_id)
+                if expansion_key not in revised_map_index_task_ids:
                     revised_tis = self._revise_map_indexes_if_mapped(
-                        schedulable.task, 
dag_version_id=schedulable.dag_version_id, session=session
+                        schedulable.task,
+                        dag_version_id=schedulable.dag_version_id,
+                        region_id=schedulable.region_id,
+                        session=session,
                     )
                     ready_tis.extend(revised_tis)
-                    revised_map_index_task_ids.add(schedulable.task.task_id)
+                    revised_map_index_task_ids.add(expansion_key)
                     if revised_tis:
                         # Revising a mapped task can add new instances, 
growing its instance count
                         # the same way expansion does. Drop the upstream-count 
memo so a downstream
@@ -1781,10 +1789,9 @@ class DagRun(Base, LoggingMixin):
                     ready_tis.append(schedulable)
 
         # Check if any ti changed state
-        tis_filter = TI.filter_for_tis(old_states)
-        if tis_filter is not None:
-            fresh_tis = session.scalars(select(TI).where(tis_filter)).all()
-            changed_tis = any(ti.state != old_states[ti.key] for ti in 
fresh_tis)
+        if old_states:
+            fresh_tis = 
session.scalars(select(TI).where(TI.id.in_(old_states))).all()
+            changed_tis = any(ti.state != old_states[ti.id] for ti in 
fresh_tis)
 
         return ready_tis, changed_tis, expansion_happened
 
@@ -2154,7 +2161,12 @@ class DagRun(Base, LoggingMixin):
             session.rollback()
 
     def _revise_map_indexes_if_mapped(
-        self, task: Operator, *, dag_version_id: UUID | None, session: Session
+        self,
+        task: Operator,
+        *,
+        dag_version_id: UUID | None,
+        session: Session,
+        region_id: UUID = SENTINEL_REGION_ID,
     ) -> list[TI]:
         """
         Check if task increased or reduced in length and handle appropriately.
@@ -2179,7 +2191,7 @@ class DagRun(Base, LoggingMixin):
                 TI.dag_id == self.dag_id,
                 TI.task_id == task.task_id,
                 TI.run_id == self.run_id,
-                TI.region_id == SENTINEL_REGION_ID,
+                TI.region_id == region_id,
             )
         )
         existing_indexes = set(query)
@@ -2192,7 +2204,7 @@ class DagRun(Base, LoggingMixin):
                     TI.dag_id == self.dag_id,
                     TI.task_id == task.task_id,
                     TI.run_id == self.run_id,
-                    TI.region_id == SENTINEL_REGION_ID,
+                    TI.region_id == region_id,
                     TI.map_index.in_(removed_indexes),
                 )
                 .values(state=TaskInstanceState.REMOVED)
@@ -2207,13 +2219,20 @@ class DagRun(Base, LoggingMixin):
             task_id=task.task_id,
             run_id=self.run_id,
             map_indexes=missing_indexes,
-            region_id=SENTINEL_REGION_ID,
+            region_id=region_id,
             session=session,
         )
 
         new_tis: list[TI] = []
         for index in missing_indexes:
-            ti = TI(task, run_id=self.run_id, map_index=index, state=None, 
dag_version_id=dag_version_id)
+            ti = TI(
+                task,
+                run_id=self.run_id,
+                map_index=index,
+                region_id=region_id,
+                state=None,
+                dag_version_id=dag_version_id,
+            )
             ti.try_number = last_tries.get(index, -1) + 1
             ti.max_tries += ti.try_number
             self.log.debug("Expanding TIs upserted %s", ti)
diff --git a/airflow-core/src/airflow/models/dynamic_region.py 
b/airflow-core/src/airflow/models/dynamic_region.py
index b453eaf031a..72e02449fc4 100644
--- a/airflow-core/src/airflow/models/dynamic_region.py
+++ b/airflow-core/src/airflow/models/dynamic_region.py
@@ -16,11 +16,14 @@
 # under the License.
 from __future__ import annotations
 
+from collections.abc import Collection
 from datetime import datetime
+from typing import TYPE_CHECKING
 from uuid import UUID
 
+import attrs
 import uuid6
-from sqlalchemy import CheckConstraint, ForeignKeyConstraint, Index, Integer, 
UniqueConstraint, Uuid
+from sqlalchemy import CheckConstraint, ForeignKeyConstraint, Index, Integer, 
UniqueConstraint, Uuid, select
 from sqlalchemy.orm import Mapped, mapped_column
 
 from airflow._shared.timezones import timezone
@@ -29,6 +32,25 @@ from airflow.utils.sqlalchemy import UtcDateTime
 
 SENTINEL_REGION_ID = UUID(int=0)
 
+if TYPE_CHECKING:
+    from sqlalchemy.orm import Session
+
+    from airflow.models.taskinstance import TaskInstance
+
+
[email protected](frozen=True)
+class ProducerContext:
+    """Caller coordinates and the producer's shared loop context from the 
pinned graph."""
+
+    region_id: UUID
+    region_index: int
+    loop_node_id: str | None = None
+    previous_iteration: bool = False
+
+
+class AmbiguousProducerError(ValueError):
+    """Multiple live executions occupy the requested producer slot."""
+
 
 class DynamicRegion(Base):
     """
@@ -78,3 +100,118 @@ class DynamicRegion(Base):
         Index("idx_dynamic_region_slot", dag_id, run_id, node_id, 
parent_region_id, parent_region_index),
         Index("idx_dynamic_region_parent_region_id", parent_region_id),
     )
+
+
+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
+
+    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")
+    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)
+    candidates = session.scalars(query).all()
+    zero = SENTINEL_REGION_ID
+    regions: dict[UUID, DynamicRegion] = {}
+    pending = ({ti.region_id for ti in candidates} - {zero}) if region_id is 
None else set()
+    if context and region_id is None:
+        pending.add(context.region_id)
+        pending.discard(zero)
+    while pending:
+        rows = session.scalars(
+            select(DynamicRegion).where(
+                DynamicRegion.dag_id == dag_id,
+                DynamicRegion.run_id == run_id,
+                DynamicRegion.id.in_(pending),
+            )
+        ).all()
+        found = {row.id for row in rows}
+        if found != pending:
+            raise ValueError("Region context does not belong to the requested 
DagRun")
+        regions.update((row.id, row) for row in rows)
+        pending = {
+            ref
+            for row in rows
+            for ref in (row.parent_region_id, row.forked_from_region_id)
+            if ref is not None and ref not in regions
+        }
+
+    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
+
+    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 ()
+
+    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))
diff --git a/airflow-core/src/airflow/models/xcom.py 
b/airflow-core/src/airflow/models/xcom.py
index e64789cb8ce..75c85b4dd83 100644
--- a/airflow-core/src/airflow/models/xcom.py
+++ b/airflow-core/src/airflow/models/xcom.py
@@ -51,6 +51,7 @@ from sqlalchemy.sql.visitors import cloned_traverse
 
 from airflow._shared.timezones import timezone
 from airflow.models.base import COLLATION_ARGS, ID_LEN, Base, 
TaskInstanceDependencies
+from airflow.models.dynamic_region import SENTINEL_REGION_ID
 from airflow.utils.db import LazySelectSequence
 from airflow.utils.helpers import is_container
 from airflow.utils.json import XComDecoder, XComEncoder
@@ -315,6 +316,8 @@ class _XComOperations:
         task_ids: str | Iterable[str] | None = None,
         dag_ids: str | Iterable[str] | None = None,
         map_indexes: int | Iterable[int] | None = None,
+        region_id: UUID | None = SENTINEL_REGION_ID,
+        producer_ids: Select | None = None,
         include_prior_dates: bool = False,
         limit: int | None = None,
         try_number: int | None = None,
@@ -325,6 +328,10 @@ class _XComOperations:
         This function returns an SQLAlchemy query of full XCom objects. If you
         just want one stored value, use :meth:`get_one` instead.
 
+        ``region_id`` is the exact producer region (the legacy sentinel by 
default); pass ``None`` to
+        enumerate across regions. ``producer_ids`` replaces the coordinate 
filters with attempts already
+        resolved by 
:func:`~airflow.models.dynamic_region.resolve_current_producers`.
+
         Use :func:`xcom_entity` for columns added to the returned statement.
 
         :param run_id: DAG run ID for the task.
@@ -347,17 +354,21 @@ class _XComOperations:
             raise ValueError(f"XCom key must be a non-empty string. Received: 
{key!r}")
         if not run_id:
             raise ValueError(f"run_id must be passed. Passed run_id={run_id}")
-        statement = build_xcom_read_query(
-            producer_ids=select_producers(
+        if producer_ids is None:
+            if include_prior_dates and region_id not in (None, 
SENTINEL_REGION_ID):
+                raise ValueError(
+                    "Prior-run lookup requires producer coordinates resolved 
separately for each run"
+                )
+            producer_ids = select_producers(
                 run_id=run_id,
                 task_ids=task_ids,
                 dag_ids=dag_ids,
                 map_indexes=map_indexes,
+                region_id=region_id,
                 include_prior_dates=include_prior_dates,
                 try_number=try_number,
-            ),
-            key=key,
-        )
+            )
+        statement = build_xcom_read_query(producer_ids=producer_ids, key=key)
         entity = xcom_entity(statement)
         statement = statement.order_by(entity.logical_date.desc(), 
entity.timestamp.desc())
         if limit:
@@ -578,6 +589,7 @@ def select_producers(
     dag_ids=None,
     task_ids=None,
     map_indexes=None,
+    region_id=SENTINEL_REGION_ID,
     include_prior_dates=False,
     try_number=None,
 ):
@@ -585,6 +597,8 @@ def select_producers(
     from airflow.models.taskinstance import TaskInstance
 
     query = select(TaskInstance.id)
+    if region_id is not None:
+        query = query.where(TaskInstance.region_id == region_id)
     if try_number is not None:
         query = query.where(TaskInstance.try_number == try_number)
     for column, value in ((TaskInstance.dag_id, dag_ids), 
(TaskInstance.task_id, task_ids)):
diff --git a/airflow-core/src/airflow/serialization/definitions/xcom_arg.py 
b/airflow-core/src/airflow/serialization/definitions/xcom_arg.py
index 748114cb810..991ab48618b 100644
--- a/airflow-core/src/airflow/serialization/definitions/xcom_arg.py
+++ b/airflow-core/src/airflow/serialization/definitions/xcom_arg.py
@@ -17,14 +17,15 @@
 
 from __future__ import annotations
 
-from collections.abc import Iterator, Sequence
+from collections.abc import Iterator, Mapping, Sequence
 from functools import singledispatch
 from typing import TYPE_CHECKING, Any
 
 import attrs
-from sqlalchemy import func, or_
+from sqlalchemy import func, or_, select
 from sqlalchemy.orm import Session
 
+from airflow.models.dynamic_region import SENTINEL_REGION_ID, ProducerContext, 
resolve_current_producers
 from airflow.models.referencemixin import ReferenceMixin
 from airflow.models.xcom import XCOM_RETURN_KEY
 from airflow.serialization.definitions.notset import NOTSET, is_arg_set
@@ -146,26 +147,60 @@ class SchedulerZipXComArg(SchedulerXComArg):
 
 
 @singledispatch
-def get_task_map_length(xcom_arg: SchedulerXComArg, run_id: str, *, session: 
Session) -> int | None:
+def get_task_map_length(
+    xcom_arg: SchedulerXComArg,
+    run_id: str,
+    *,
+    producer_contexts: Mapping[str, ProducerContext] | None = None,
+    session: Session,
+) -> int | None:
     # The base implementation -- specific XComArg subclasses have specialised 
implementations
     raise NotImplementedError(f"get_task_map_length not implemented for 
{type(xcom_arg)}")
 
 
 @get_task_map_length.register
-def _(xcom_arg: SchedulerPlainXComArg, run_id: str, *, session: Session) -> 
int | None:
+def _(
+    xcom_arg: SchedulerPlainXComArg,
+    run_id: str,
+    *,
+    producer_contexts: Mapping[str, ProducerContext] | None = None,
+    session: Session,
+) -> int | None:
     from airflow.models.taskinstance import TaskInstance
     from airflow.models.xcom import XComModel, xcom_entity
     from airflow.serialization.definitions.mappedoperator import is_mapped
 
     dag_id = xcom_arg.operator.dag_id
     task_id = xcom_arg.operator.task_id
-
-    if is_mapped(xcom_arg.operator):
+    mapped = is_mapped(xcom_arg.operator)
+
+    if producer_contexts is not None:
+        producers = resolve_current_producers(
+            dag_id=dag_id,
+            run_id=run_id,
+            task_id=task_id,
+            is_mapped=mapped,
+            context=producer_contexts.get(task_id),
+            session=session,
+        )
+        if not producers:
+            return None
+        if mapped and any(ti.state in State.unfinished for ti in producers):
+            return None
+        read = XComModel.get_many(
+            dag_ids=dag_id,
+            run_id=run_id,
+            task_ids=task_id,
+            key=XCOM_RETURN_KEY,
+            
producer_ids=select(TaskInstance.id).where(TaskInstance.id.in_([ti.id for ti in 
producers])),
+        )
+    elif mapped:
         unfinished_ti_exists = exists_query(
             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,
             # Special NULL treatment is needed because 'state' can be NULL.
             # The "IN" part would produce "NULL NOT IN ..." and eventually
             # "NULl = NULL", which is a big no-no in SQL.
@@ -178,26 +213,45 @@ def _(xcom_arg: SchedulerPlainXComArg, run_id: str, *, 
session: Session) -> int
         if unfinished_ti_exists:
             return None  # Not all of the expanded tis are done yet.
         read = XComModel.get_many(dag_ids=dag_id, run_id=run_id, 
task_ids=task_id, key=XCOM_RETURN_KEY)
-        entity = xcom_entity(read)
-        return session.scalar(
-            read.order_by(None).where(entity.map_index >= 
0).with_only_columns(func.count(entity.map_index))
+    else:
+        read = XComModel.get_many(
+            dag_ids=dag_id, run_id=run_id, task_ids=task_id, map_indexes=-1, 
key=XCOM_RETURN_KEY
         )
 
-    read = XComModel.get_many(
-        dag_ids=dag_id, run_id=run_id, task_ids=task_id, map_indexes=-1, 
key=XCOM_RETURN_KEY
-    )
     entity = xcom_entity(read)
-    return session.scalar(read.with_only_columns(entity.mapped_length))
+    if mapped:
+        return session.scalar(
+            read.order_by(None).where(entity.map_index >= 
0).with_only_columns(func.count(entity.map_index))
+        )
+    # Not xcom_arg.key: the SDK records the length of the whole return value, 
never per key.
+    if producer_contexts is None:
+        read = read.where(entity.map_index == -1)
+    return 
session.scalar(read.order_by(None).with_only_columns(entity.mapped_length))
 
 
 @get_task_map_length.register
-def _(xcom_arg: SchedulerMapXComArg, run_id: str, *, session: Session) -> int 
| None:
-    return get_task_map_length(xcom_arg.arg, run_id, session=session)
+def _(
+    xcom_arg: SchedulerMapXComArg,
+    run_id: str,
+    *,
+    producer_contexts: Mapping[str, ProducerContext] | None = None,
+    session: Session,
+) -> int | None:
+    return get_task_map_length(xcom_arg.arg, run_id, 
producer_contexts=producer_contexts, session=session)
 
 
 @get_task_map_length.register
-def _(xcom_arg: SchedulerZipXComArg, run_id: str, *, session: Session) -> int 
| None:
-    all_lengths = (get_task_map_length(arg, run_id, session=session) for arg 
in xcom_arg.args)
+def _(
+    xcom_arg: SchedulerZipXComArg,
+    run_id: str,
+    *,
+    producer_contexts: Mapping[str, ProducerContext] | None = None,
+    session: Session,
+) -> int | None:
+    all_lengths = (
+        get_task_map_length(arg, run_id, producer_contexts=producer_contexts, 
session=session)
+        for arg in xcom_arg.args
+    )
     ready_lengths = [length for length in all_lengths if length is not None]
     if len(ready_lengths) != len(xcom_arg.args):
         return None  # If any of the referenced XComs is not ready, we are not 
ready either.
@@ -207,8 +261,17 @@ def _(xcom_arg: SchedulerZipXComArg, run_id: str, *, 
session: Session) -> int |
 
 
 @get_task_map_length.register
-def _(xcom_arg: SchedulerConcatXComArg, run_id: str, *, session: Session) -> 
int | None:
-    all_lengths = (get_task_map_length(arg, run_id, session=session) for arg 
in xcom_arg.args)
+def _(
+    xcom_arg: SchedulerConcatXComArg,
+    run_id: str,
+    *,
+    producer_contexts: Mapping[str, ProducerContext] | None = None,
+    session: Session,
+) -> int | None:
+    all_lengths = (
+        get_task_map_length(arg, run_id, producer_contexts=producer_contexts, 
session=session)
+        for arg in xcom_arg.args
+    )
     ready_lengths = [length for length in all_lengths if length is not None]
     if len(ready_lengths) != len(xcom_arg.args):
         return None  # If any of the referenced XComs is not ready, we are not 
ready either.
diff --git a/airflow-core/tests/unit/models/test_dynamic_region.py 
b/airflow-core/tests/unit/models/test_dynamic_region.py
new file mode 100644
index 00000000000..2859592944a
--- /dev/null
+++ b/airflow-core/tests/unit/models/test_dynamic_region.py
@@ -0,0 +1,378 @@
+# 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
+from uuid import uuid4
+
+import pytest
+from sqlalchemy import select
+
+from airflow._shared.timezones import timezone
+from airflow.models.dynamic_region import (
+    AmbiguousProducerError,
+    DynamicRegion,
+    ProducerContext,
+    resolve_current_producers,
+)
+from airflow.models.taskinstance import TaskInstance
+from airflow.models.xcom import XComModel
+from airflow.providers.standard.operators.empty import EmptyOperator
+from airflow.providers.standard.operators.python import PythonOperator
+from airflow.utils.state import TaskInstanceState
+
+from tests_common.test_utils.db import clear_db_runs
+
+if TYPE_CHECKING:
+    from airflow.models.dagrun import DagRun
+
+pytestmark = pytest.mark.db_test
+
+
[email protected](autouse=True)
+def clean_db():
+    clear_db_runs()
+    yield
+    clear_db_runs()
+
+
[email protected]
+def dag_run(dag_maker):
+    with dag_maker(serialized=True):
+        EmptyOperator(task_id="task")
+    return dag_maker.create_dagrun()
+
+
+def make_region(dag_run: DagRun, **kwargs) -> DynamicRegion:
+    return DynamicRegion(dag_id=dag_run.dag_id, run_id=dag_run.run_id, 
node_id="loop", **kwargs)
+
+
[email protected]
+def regional_tis(dag_maker, session):
+    with dag_maker(serialized=True):
+        task = EmptyOperator(task_id="task")
+    dr = dag_maker.create_dagrun()
+    original = dr.task_instances[0]
+    other = TaskInstance(task=task, run_id=dr.run_id, 
dag_version_id=original.dag_version_id)
+    other.region_id = uuid4()
+    session.add(other)
+    session.flush()
+    return original, other
+
+
[email protected]
+def producer_tis(dag_maker, session):
+    with dag_maker(serialized=True) as dag:
+        for task_id in ("producer", "consumer", "outside", "mapped"):
+            EmptyOperator(task_id=task_id)
+    dr = dag_maker.create_dagrun()
+    regions = []
+    for _ in range(3):
+        region = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id, 
node_id="loop")
+        if regions:
+            region.forked_from_region_id = regions[-1].id
+        session.add(region)
+        session.flush()
+        regions.append(region)
+    tis = {ti.task_id: ti for ti in dr.task_instances}
+    tis["producer"].region_id = regions[0].id
+    tis["producer"].region_index = 2
+    tis["consumer"].region_id = regions[2].id
+    tis["consumer"].region_index = 2
+    previous = TaskInstance(
+        dag.get_task("producer"),
+        tis["producer"].dag_version_id,
+        run_id=dr.run_id,
+        map_index=1,
+        region_id=regions[0].id,
+    )
+    session.add(previous)
+    session.flush()
+    return tis, regions, previous
+
+
[email protected]("previous_iteration", [False, True])
+def test_resolve_retained_producer_across_repeated_forks(producer_tis, 
session, previous_iteration):
+    tis, regions, previous = producer_tis
+    consumer = tis["consumer"]
+    selected = resolve_current_producers(
+        dag_id=consumer.dag_id,
+        run_id=consumer.run_id,
+        task_id="producer",
+        is_mapped=False,
+        context=ProducerContext(consumer.region_id, consumer.region_index, 
"loop", previous_iteration),
+        session=session,
+    )
+    assert [ti.id for ti in selected] == [previous.id if previous_iteration 
else tis["producer"].id]
+
+
+def test_resolve_producer_outside_the_loop(producer_tis, session):
+    tis, _, _ = producer_tis
+    consumer = tis["consumer"]
+    selected = resolve_current_producers(
+        dag_id=consumer.dag_id,
+        run_id=consumer.run_id,
+        task_id="outside",
+        is_mapped=False,
+        context=ProducerContext(consumer.region_id, consumer.region_index),
+        session=session,
+    )
+    assert [ti.id for ti in selected] == [tis["outside"].id]
+
+
+def test_resolver_never_revives_archived_producer(producer_tis, session):
+    tis, _, _ = producer_tis
+    producer, consumer = tis["producer"], tis["consumer"]
+    producer.archive(reason="test", session=session)
+    assert (
+        resolve_current_producers(
+            dag_id=consumer.dag_id,
+            run_id=consumer.run_id,
+            task_id="producer",
+            is_mapped=False,
+            context=ProducerContext(consumer.region_id, consumer.region_index, 
"loop"),
+            session=session,
+        )
+        == ()
+    )
+
+
+def test_resolver_rejects_ambiguous_live_producers(producer_tis, session):
+    tis, regions, _ = producer_tis
+    producer, consumer = tis["producer"], tis["consumer"]
+    other = TaskInstance(
+        producer.task,
+        producer.dag_version_id,
+        run_id=producer.run_id,
+        map_index=producer.map_index,
+        region_id=regions[1].id,
+    )
+    session.add(other)
+    session.flush()
+    with pytest.raises(AmbiguousProducerError):
+        resolve_current_producers(
+            dag_id=consumer.dag_id,
+            run_id=consumer.run_id,
+            task_id="producer",
+            is_mapped=False,
+            context=ProducerContext(consumer.region_id, consumer.region_index, 
"loop"),
+            session=session,
+        )
+
+
[email protected]("mapped_caller", [False, True])
+def test_mapped_producer_scope_precedes_index_selection(producer_tis, session, 
mapped_caller):
+    tis, regions, _ = producer_tis
+    consumer, mapped = tis["consumer"], tis["mapped"]
+    children = []
+    for parent, iteration in ((regions[0], 2), (regions[2], 2), (regions[2], 
3)):
+        child = DynamicRegion(
+            dag_id=mapped.dag_id,
+            run_id=mapped.run_id,
+            node_id="mapped",
+            parent_region_id=parent.id,
+            parent_region_index=iteration,
+        )
+        session.add(child)
+        session.flush()
+        children.append(child)
+    mapped.region_id, mapped.region_index = children[0].id, 0
+    second = TaskInstance(
+        mapped.task, mapped.dag_version_id, run_id=mapped.run_id, map_index=1, 
region_id=children[1].id
+    )
+    wrong_iteration = TaskInstance(
+        mapped.task, mapped.dag_version_id, run_id=mapped.run_id, map_index=0, 
region_id=children[2].id
+    )
+    session.add_all([second, wrong_iteration])
+    if mapped_caller:
+        caller_region = DynamicRegion(
+            dag_id=consumer.dag_id,
+            run_id=consumer.run_id,
+            node_id="consumer",
+            parent_region_id=regions[2].id,
+            parent_region_index=2,
+        )
+        session.add(caller_region)
+        session.flush()
+        consumer.region_id, consumer.region_index = caller_region.id, 5
+    session.flush()
+    context = ProducerContext(consumer.region_id, consumer.region_index, 
"loop")
+    selected = resolve_current_producers(
+        dag_id=mapped.dag_id,
+        run_id=mapped.run_id,
+        task_id="mapped",
+        is_mapped=True,
+        context=context,
+        session=session,
+    )
+    assert [ti.id for ti in selected] == [mapped.id, second.id]
+    selected = resolve_current_producers(
+        dag_id=mapped.dag_id,
+        run_id=mapped.run_id,
+        task_id="mapped",
+        is_mapped=True,
+        context=context,
+        map_indexes=1,
+        session=session,
+    )
+    assert [ti.id for ti in selected] == [second.id]
+
+
+def test_previous_iteration_zero_is_missing(producer_tis, session):
+    tis, _, _ = producer_tis
+    consumer = tis["consumer"]
+    assert (
+        resolve_current_producers(
+            dag_id=consumer.dag_id,
+            run_id=consumer.run_id,
+            task_id="producer",
+            is_mapped=False,
+            context=ProducerContext(consumer.region_id, 0, "loop", True),
+            session=session,
+        )
+        == ()
+    )
+
+
+def test_explicit_producer_coordinate_is_task_and_run_scoped(regional_tis, 
session):
+    first, second = regional_tis
+    selected = resolve_current_producers(
+        dag_id=second.dag_id,
+        run_id=second.run_id,
+        task_id=second.task_id,
+        is_mapped=False,
+        region_id=second.region_id,
+        region_index=second.region_index,
+        session=session,
+    )
+    assert [ti.id for ti in selected] == [second.id]
+    assert (
+        resolve_current_producers(
+            dag_id=first.dag_id,
+            run_id="missing",
+            task_id=first.task_id,
+            is_mapped=False,
+            region_id=second.region_id,
+            region_index=second.region_index,
+            session=session,
+        )
+        == ()
+    )
+
+
+def test_region_exact_ti_lookup(regional_tis, session):
+    first, second = regional_tis
+    found = TaskInstance.get_task_instance(
+        second.dag_id,
+        second.run_id,
+        second.task_id,
+        second.map_index,
+        region_id=second.region_id,
+        session=session,
+    )
+    assert found.id == second.id
+    assert (
+        second.dag_run.get_task_instance(second.task_id, 
region_id=second.region_id, session=session).id
+        == second.id
+    )
+    assert 
session.scalar(select(TaskInstance).where(TaskInstance.filter_for_tis([second]))).id
 == second.id
+    assert 
session.scalar(select(TaskInstance).where(TaskInstance.filter_for_tis([first.key]))).id
 == first.id
+
+
+def test_dependency_state_change_is_correlated_by_uuid(regional_tis, session, 
mocker):
+    first, second = regional_tis
+
+    def fail_dependency(ti, **kwargs):
+        if ti.id == second.id:
+            ti.state = TaskInstanceState.UPSTREAM_FAILED
+            session.flush()
+        return False
+
+    mocker.patch.object(TaskInstance, "are_dependencies_met", autospec=True, 
side_effect=fail_dependency)
+    ready, changed, expanded = 
first.dag_run._get_ready_tis(list(regional_tis), [], session=session)
+    assert ready == []
+    assert changed is True
+    assert expanded is False
+    assert first.state is None
+    assert second.state == TaskInstanceState.UPSTREAM_FAILED
+
+
+def test_mapping_revision_only_changes_selected_expansion(dag_maker, session):
+    with dag_maker(serialized=True):
+        PythonOperator.partial(task_id="mapped", python_callable=lambda: 
None).expand(op_kwargs=[{}, {}])
+    dr = dag_maker.create_dagrun()
+    task = dr.get_dag().get_task("mapped")
+    version = dr.task_instances[0].dag_version_id
+    ordinary = TaskInstance(task, version, run_id=dr.run_id, map_index=3)
+    regional = TaskInstance(task, version, run_id=dr.run_id, map_index=3, 
region_id=uuid4())
+    session.add_all([ordinary, regional])
+    session.flush()
+    added = dr._revise_map_indexes_if_mapped(
+        task, dag_version_id=version, region_id=regional.region_id, 
session=session
+    )
+    assert [(ti.region_id, ti.region_index) for ti in added] == [
+        (regional.region_id, 0),
+        (regional.region_id, 1),
+    ]
+    session.refresh(regional)
+    session.refresh(ordinary)
+    assert regional.state == TaskInstanceState.REMOVED
+    assert ordinary.state is None
+
+
+def test_prior_run_xcom_uses_resolved_producers_for_each_run(dag_maker, 
session):
+    with dag_maker(serialized=True):
+        EmptyOperator(task_id="task")
+    tis = []
+    for day in (1, 2):
+        ti = dag_maker.create_dagrun(
+            run_id=f"run-{day}", logical_date=timezone.datetime(2026, 1, day)
+        ).task_instances[0]
+        ti.region_id = uuid4()
+        ti.region_index = 0
+        session.flush()
+        XComModel.set_for_attempt(task_instance_id=ti.id, key="key", 
value=day, session=session)
+        tis.append(ti)
+    rows = session.scalars(
+        XComModel.get_many(
+            run_id=tis[1].run_id,
+            dag_ids=tis[1].dag_id,
+            task_ids="task",
+            key="key",
+            include_prior_dates=True,
+            
producer_ids=select(TaskInstance.id).where(TaskInstance.id.in_([ti.id for ti in 
tis])),
+        )
+    ).all()
+    assert [row.task_instance_id for row in rows] == [tis[1].id, tis[0].id]
+    with pytest.raises(ValueError, match="resolved separately"):
+        XComModel.get_many(run_id=tis[1].run_id, region_id=tis[1].region_id, 
include_prior_dates=True)
+
+
+def test_xcom_reads_are_scoped_to_the_producer_region(regional_tis, session):
+    first, second = regional_tis
+    for ti, value in ((first, "ordinary"), (second, "regional")):
+        XComModel.set_for_attempt(task_instance_id=ti.id, key="key", 
value=value, session=session)
+    session.flush()
+
+    def read(**kwargs):
+        statement = XComModel.get_many(run_id=first.run_id, 
dag_ids=first.dag_id, key="key", **kwargs)
+        return {row.task_instance_id for row in session.scalars(statement)}
+
+    assert read() == {first.id}
+    assert read(region_id=second.region_id) == {second.id}
+    assert read(region_id=None) == {first.id, second.id}
diff --git a/airflow-core/tests/unit/models/test_xcom_arg.py 
b/airflow-core/tests/unit/models/test_xcom_arg.py
index 21f13c50710..043cb5a050f 100644
--- a/airflow-core/tests/unit/models/test_xcom_arg.py
+++ b/airflow-core/tests/unit/models/test_xcom_arg.py
@@ -18,13 +18,23 @@ from __future__ import annotations
 
 import pytest
 
+from airflow.models.dynamic_region import DynamicRegion, ProducerContext
 from airflow.models.expandinput import NotFullyPopulated
+from airflow.models.taskinstance import TaskInstance
 from airflow.models.xcom import XCOM_RETURN_KEY, XComModel
 from airflow.models.xcom_arg import XComArg
 from airflow.providers.standard.operators.bash import BashOperator
 from airflow.providers.standard.operators.python import PythonOperator
 from airflow.serialization.definitions.mappedoperator import 
get_mapped_ti_count
 from airflow.serialization.definitions.notset import NOTSET
+from airflow.serialization.definitions.xcom_arg import (
+    SchedulerConcatXComArg,
+    SchedulerMapXComArg,
+    SchedulerPlainXComArg,
+    SchedulerZipXComArg,
+    get_task_map_length,
+)
+from airflow.utils.state import TaskInstanceState
 
 from tests_common.test_utils.db import clear_db_dags, clear_db_runs
 
@@ -266,3 +276,144 @@ def 
test_mapped_length_dies_with_the_pushed_value(dag_maker, session):
         session=session,
     )
     assert get_mapped_ti_count(consume_task, dr.run_id, session=session) == 2
+
+
[email protected](
+    ("operation", "expected_length"),
+    [("plain", 2), ("map", 2), ("zip", 2), ("zip_longest", 5), ("concat", 7)],
+)
+def test_map_length_selects_retained_producer_in_callers_iteration(
+    dag_maker, session, operation, expected_length
+):
+    with dag_maker(session=session, serialized=True) as dag:
+
+        @dag.task
+        def source():
+            return [1, 2]
+
+        @dag.task
+        def outside():
+            return [1, 2, 3, 4, 5]
+
+        source()
+        outside()
+
+    dr = dag_maker.create_dagrun()
+    original = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id, 
node_id="loop")
+    session.add(original)
+    session.flush()
+    replacement = DynamicRegion(
+        dag_id=dr.dag_id,
+        run_id=dr.run_id,
+        node_id="loop",
+        forked_from_region_id=original.id,
+        resumes_from_index=1,
+    )
+    session.add(replacement)
+    source_ti = next(ti for ti in dr.task_instances if ti.task_id == "source")
+    outside_ti = next(ti for ti in dr.task_instances if ti.task_id == 
"outside")
+    source_ti.region_id = original.id
+    source_ti.region_index = 1
+    source_ti.state = TaskInstanceState.SUCCESS
+    earlier_ti = TaskInstance(
+        task=dag_maker.serialized_dag.get_task("source"),
+        run_id=dr.run_id,
+        map_index=0,
+        dag_version_id=source_ti.dag_version_id,
+    )
+    earlier_ti.region_id = original.id
+    earlier_ti.state = TaskInstanceState.SUCCESS
+    session.add(earlier_ti)
+    session.flush()
+    for ti, length in ((earlier_ti, 99), (source_ti, 2), (outside_ti, 5)):
+        XComModel.set_for_attempt(
+            task_instance_id=ti.id,
+            key=XCOM_RETURN_KEY,
+            value=list(range(length)),
+            mapped_length=length,
+            session=session,
+        )
+    session.flush()
+    source_arg = 
SchedulerPlainXComArg(dag_maker.serialized_dag.get_task("source"), 
XCOM_RETURN_KEY)
+    outside_arg = 
SchedulerPlainXComArg(dag_maker.serialized_dag.get_task("outside"), 
XCOM_RETURN_KEY)
+    argument = {
+        "plain": source_arg,
+        "map": SchedulerMapXComArg(source_arg, ["str"]),
+        "zip": SchedulerZipXComArg([source_arg, outside_arg], NOTSET),
+        "zip_longest": SchedulerZipXComArg([source_arg, outside_arg], None),
+        "concat": SchedulerConcatXComArg([source_arg, outside_arg]),
+    }[operation]
+    contexts = {"source": ProducerContext(replacement.id, 1, 
loop_node_id="loop")}
+
+    assert (
+        get_task_map_length(argument, dr.run_id, producer_contexts=contexts, 
session=session)
+        == expected_length
+    )
+
+
[email protected](
+    ("count", "unfinished", "expected_length"), [(0, False, 0), (2, False, 2), 
(2, True, None)]
+)
+def test_mapped_producer_length_ignores_other_iterations(
+    dag_maker, session, count, unfinished, expected_length
+):
+    with dag_maker(session=session, serialized=True) as dag:
+
+        @dag.task
+        def source(value):
+            return value
+
+        source.expand(value=list(range(count)))
+
+    dr = dag_maker.create_dagrun()
+    loop = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id, node_id="loop")
+    session.add(loop)
+    session.flush()
+    expansions = [
+        DynamicRegion(
+            dag_id=dr.dag_id,
+            run_id=dr.run_id,
+            node_id="source",
+            parent_region_id=loop.id,
+            parent_region_index=index,
+        )
+        for index in range(2)
+    ]
+    session.add_all(expansions)
+    session.flush()
+    for ti in dr.task_instances:
+        ti.region_id = expansions[1].id
+        ti.state = TaskInstanceState.SUCCESS if count else 
TaskInstanceState.SKIPPED
+    session.flush()
+    for ti in dr.task_instances:
+        if ti.map_index >= 0:
+            XComModel.set_for_attempt(
+                task_instance_id=ti.id, key=XCOM_RETURN_KEY, 
value=ti.map_index, session=session
+            )
+    if unfinished:
+        dr.task_instances[0].state = TaskInstanceState.RUNNING
+    earlier_ti = TaskInstance(
+        task=dag_maker.serialized_dag.get_task("source"),
+        run_id=dr.run_id,
+        map_index=0,
+        dag_version_id=dr.task_instances[0].dag_version_id,
+    )
+    earlier_ti.region_id = expansions[0].id
+    earlier_ti.state = TaskInstanceState.RUNNING
+    session.add(earlier_ti)
+    session.flush()
+    XComModel.set_for_attempt(
+        task_instance_id=earlier_ti.id, key=XCOM_RETURN_KEY, value="other 
iteration", session=session
+    )
+    session.flush()
+    argument = 
SchedulerPlainXComArg(dag_maker.serialized_dag.get_task("source"), 
XCOM_RETURN_KEY)
+
+    assert (
+        get_task_map_length(
+            argument,
+            dr.run_id,
+            producer_contexts={"source": ProducerContext(loop.id, 1, 
loop_node_id="loop")},
+            session=session,
+        )
+        == expected_length
+    )

Reply via email to