Skip to content

Commit 8d326be

Browse files
committed
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).
1 parent 049690d commit 8d326be

6 files changed

Lines changed: 1011 additions & 43 deletions

File tree

‎airflow-core/src/airflow/models/dagrun.py‎

Lines changed: 36 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -124,7 +124,6 @@
124124
TaskInstance as TIDataModel,
125125
)
126126
from airflow.models.dag_version import DagVersion
127-
from airflow.models.taskinstancekey import TaskInstanceKey
128127
from airflow.sdk import DAG as SDKDAG
129128
from airflow.serialization.definitions.dag import SerializedDAG
130129
from airflow.serialization.definitions.mappedoperator import Operator
@@ -1064,6 +1063,7 @@ def get_task_instance(
10641063
task_id: str,
10651064
*,
10661065
map_index: int = -1,
1066+
region_id: UUID = SENTINEL_REGION_ID,
10671067
session: Session = NEW_SESSION,
10681068
) -> TI | None:
10691069
"""
@@ -1078,6 +1078,7 @@ def get_task_instance(
10781078
task_id=task_id,
10791079
session=session,
10801080
map_index=map_index,
1081+
region_id=region_id,
10811082
)
10821083

10831084
@staticmethod
@@ -1088,6 +1089,7 @@ def fetch_task_instance(
10881089
task_id: str,
10891090
*,
10901091
map_index: int = -1,
1092+
region_id: UUID = SENTINEL_REGION_ID,
10911093
session: Session = NEW_SESSION,
10921094
) -> TI | None:
10931095
"""
@@ -1099,7 +1101,9 @@ def fetch_task_instance(
10991101
:param session: Sqlalchemy ORM Session
11001102
"""
11011103
return session.scalars(
1102-
select(TI).filter_by(dag_id=dag_id, run_id=dag_run_id, task_id=task_id, map_index=map_index)
1104+
select(TI).filter_by(
1105+
dag_id=dag_id, run_id=dag_run_id, task_id=task_id, map_index=map_index, region_id=region_id
1106+
)
11031107
).one_or_none()
11041108

11051109
def get_dag(self) -> SerializedDAG:
@@ -1683,7 +1687,7 @@ def _get_ready_tis(
16831687
finished_tis: list[TI],
16841688
session: Session,
16851689
) -> tuple[list[TI], bool, bool]:
1686-
old_states: dict[TaskInstanceKey, Any] = {}
1690+
old_states: dict[UUID, Any] = {}
16871691
ready_tis: list[TI] = []
16881692
changed_tis = False
16891693

@@ -1735,13 +1739,13 @@ def _expand_mapped_task_if_needed(ti: TI) -> Iterable[TI] | None:
17351739
# Check dependencies.
17361740
expansion_happened = False
17371741
# Set of task ids for which was already done _revise_map_indexes_if_mapped
1738-
revised_map_index_task_ids: set[str] = set()
1742+
revised_map_index_task_ids: set[tuple[str, UUID]] = set()
17391743
for schedulable in itertools.chain(schedulable_tis, additional_tis):
17401744
if TYPE_CHECKING:
17411745
assert isinstance(schedulable.task, Operator)
17421746
old_state = schedulable.state
17431747
if not schedulable.are_dependencies_met(session=session, dep_context=dep_context):
1744-
old_states[schedulable.key] = old_state
1748+
old_states[schedulable.id] = old_state
17451749
continue
17461750
# If schedulable is not yet expanded, try doing it now. This is
17471751
# called in two places: First and ideally in the mini scheduler at
@@ -1762,12 +1766,16 @@ def _expand_mapped_task_if_needed(ti: TI) -> Iterable[TI] | None:
17621766
if new_tis is None and schedulable.state in SCHEDULEABLE_STATES:
17631767
# It's enough to revise map index once per task id,
17641768
# checking the map index for each mapped task significantly slows down scheduling
1765-
if schedulable.task.task_id not in revised_map_index_task_ids:
1769+
expansion_key = (schedulable.task.task_id, schedulable.region_id)
1770+
if expansion_key not in revised_map_index_task_ids:
17661771
revised_tis = self._revise_map_indexes_if_mapped(
1767-
schedulable.task, dag_version_id=schedulable.dag_version_id, session=session
1772+
schedulable.task,
1773+
dag_version_id=schedulable.dag_version_id,
1774+
region_id=schedulable.region_id,
1775+
session=session,
17681776
)
17691777
ready_tis.extend(revised_tis)
1770-
revised_map_index_task_ids.add(schedulable.task.task_id)
1778+
revised_map_index_task_ids.add(expansion_key)
17711779
if revised_tis:
17721780
# Revising a mapped task can add new instances, growing its instance count
17731781
# the same way expansion does. Drop the upstream-count memo so a downstream
@@ -1781,10 +1789,9 @@ def _expand_mapped_task_if_needed(ti: TI) -> Iterable[TI] | None:
17811789
ready_tis.append(schedulable)
17821790

17831791
# Check if any ti changed state
1784-
tis_filter = TI.filter_for_tis(old_states)
1785-
if tis_filter is not None:
1786-
fresh_tis = session.scalars(select(TI).where(tis_filter)).all()
1787-
changed_tis = any(ti.state != old_states[ti.key] for ti in fresh_tis)
1792+
if old_states:
1793+
fresh_tis = session.scalars(select(TI).where(TI.id.in_(old_states))).all()
1794+
changed_tis = any(ti.state != old_states[ti.id] for ti in fresh_tis)
17881795

17891796
return ready_tis, changed_tis, expansion_happened
17901797

@@ -2154,7 +2161,12 @@ def _create_task_instances(
21542161
session.rollback()
21552162

21562163
def _revise_map_indexes_if_mapped(
2157-
self, task: Operator, *, dag_version_id: UUID | None, session: Session
2164+
self,
2165+
task: Operator,
2166+
*,
2167+
dag_version_id: UUID | None,
2168+
session: Session,
2169+
region_id: UUID = SENTINEL_REGION_ID,
21582170
) -> list[TI]:
21592171
"""
21602172
Check if task increased or reduced in length and handle appropriately.
@@ -2179,7 +2191,7 @@ def _revise_map_indexes_if_mapped(
21792191
TI.dag_id == self.dag_id,
21802192
TI.task_id == task.task_id,
21812193
TI.run_id == self.run_id,
2182-
TI.region_id == SENTINEL_REGION_ID,
2194+
TI.region_id == region_id,
21832195
)
21842196
)
21852197
existing_indexes = set(query)
@@ -2192,7 +2204,7 @@ def _revise_map_indexes_if_mapped(
21922204
TI.dag_id == self.dag_id,
21932205
TI.task_id == task.task_id,
21942206
TI.run_id == self.run_id,
2195-
TI.region_id == SENTINEL_REGION_ID,
2207+
TI.region_id == region_id,
21962208
TI.map_index.in_(removed_indexes),
21972209
)
21982210
.values(state=TaskInstanceState.REMOVED)
@@ -2207,13 +2219,20 @@ def _revise_map_indexes_if_mapped(
22072219
task_id=task.task_id,
22082220
run_id=self.run_id,
22092221
map_indexes=missing_indexes,
2210-
region_id=SENTINEL_REGION_ID,
2222+
region_id=region_id,
22112223
session=session,
22122224
)
22132225

22142226
new_tis: list[TI] = []
22152227
for index in missing_indexes:
2216-
ti = TI(task, run_id=self.run_id, map_index=index, state=None, dag_version_id=dag_version_id)
2228+
ti = TI(
2229+
task,
2230+
run_id=self.run_id,
2231+
map_index=index,
2232+
region_id=region_id,
2233+
state=None,
2234+
dag_version_id=dag_version_id,
2235+
)
22172236
ti.try_number = last_tries.get(index, -1) + 1
22182237
ti.max_tries += ti.try_number
22192238
self.log.debug("Expanding TIs upserted %s", ti)

0 commit comments

Comments
 (0)