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


##########
airflow-core/src/airflow/models/xcom.py:
##########
@@ -315,6 +316,8 @@ def get_many(
         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,

Review Comment:
   When `producer_ids` is passed, `task_ids`, `dag_ids`, `map_indexes`, 
`region_id` and `include_prior_dates` are all silently ignored (and 
`try_number` still flips `include_all_attempts`). The map-length caller already 
passes `dag_ids`/`task_ids` that do nothing, and the prior-run test passes an 
inert `include_prior_dates=True`. Would it be clearer to raise when 
`producer_ids` is combined with any coordinate filter, or have resolver callers 
use `build_xcom_read_query(producer_ids=..., key=...)` directly?



##########
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])),

Review Comment:
   With `producer_ids` set, `get_many` skips `select_producers`, so 
`include_prior_dates=True` does nothing here, and the ids are hand-built rather 
than resolved. The assertion passes with the flag off too, so it only checks 
the existing `logical_date desc` ordering. Only the `pytest.raises` half tests 
something this PR changed. Either resolve each run's producer with 
`resolve_current_producers` and read their union (dropping the flag), or rename 
the test to what it checks.



##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -1762,12 +1766,16 @@ def _expand_mapped_task_if_needed(ti: TI) -> 
Iterable[TI] | None:
             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)

Review Comment:
   Nothing drives `_get_ready_tis` with one mapped task in two regions, so 
reverting this key to `schedulable.task.task_id` keeps every test green even 
though the second region would skip its revision that pass. 
`test_mapping_revision_only_changes_selected_expansion` calls 
`_revise_map_indexes_if_mapped` directly, and the UUID test stubs 
`are_dependencies_met` to False so it never reaches here. Could you add one 
with a stale `map_index=3` instance in two regions (expansion length 2) passed 
through `_get_ready_tis`, asserting both become REMOVED and both regions get 0 
and 1?



##########
airflow-core/src/airflow/serialization/definitions/xcom_arg.py:
##########
@@ -146,26 +147,60 @@ def iter_references(self) -> Iterator[tuple[Operator, 
str]]:
 
 
 @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,

Review Comment:
   `mapped` here is `is_mapped(operator)`, so a plain task inside a mapped task 
group counts as unmapped and every one of its expanded instances gets public 
index -1. If that group sits in a loop and a task in the same iteration does 
`.expand(x=member.output)`, all the member's instances share the caller's 
iteration, so the second one raises `AmbiguousProducerError`. 
`expand_mapped_task` only catches `NotFullyPopulated`, so the scheduler logs 
"Error scheduling DAG run" on every loop and the run never moves. Before this 
change the same Dag got `None` here and failed the expansion cleanly. Could 
this branch keep that behaviour, e.g. return `None` when the operator is not 
mapped but `get_needs_expansion()` is true? The coordinate resolver later 
passes `is_mapped=task.get_needs_expansion()`, so the two callers also disagree 
about what "mapped" means for the same producer. A test with a mapped task 
group inside a loop would pin it down.



##########
airflow-core/src/airflow/models/dynamic_region.py:
##########
@@ -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()

Review Comment:
   When the lookup is scoped by loop position, this loads every live instance 
of the producer across all iterations (full ORM rows with the joined 
`dag_run`), walks their region ancestry, and only then filters by position in 
Python. For `b.expand(x=a.output)` inside a loop with `a` mapped, iteration k 
loads I x M rows per call, and the trigger-rule dep calls this once per pending 
`b` instance per scheduling pass, so 10 iterations at width 1024 is about 10M 
rows a pass. Could the position go into SQL instead? Resolve the loop's fork 
family once, then select direct members with `region_id IN family AND 
region_index = k` and nested expansions through 
`dynamic_region.parent_region_id IN family AND parent_region_index = k`. The 
SQL-side filters added later don't cover this path.



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

Review Comment:
   The `dag_run` fixture and `make_region` aren't used by any test in this file 
(and the `TYPE_CHECKING` `DagRun` import exists only for `make_region`). Could 
they go, with `make_region` added where it gets its first caller?



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