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