Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 38 additions & 19 deletions airflow-core/src/airflow/models/dagrun.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,7 +124,6 @@
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
Expand Down Expand Up @@ -1064,6 +1063,7 @@ def get_task_instance(
task_id: str,
*,
map_index: int = -1,
region_id: UUID = SENTINEL_REGION_ID,
session: Session = NEW_SESSION,
) -> TI | None:
"""
Expand All @@ -1078,6 +1078,7 @@ def get_task_instance(
task_id=task_id,
session=session,
map_index=map_index,
region_id=region_id,
)

@staticmethod
Expand All @@ -1088,6 +1089,7 @@ def fetch_task_instance(
task_id: str,
*,
map_index: int = -1,
region_id: UUID = SENTINEL_REGION_ID,
session: Session = NEW_SESSION,
) -> TI | None:
"""
Expand All @@ -1099,7 +1101,9 @@ def fetch_task_instance(
: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:
Expand Down Expand Up @@ -1683,7 +1687,7 @@ def _get_ready_tis(
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

Expand Down Expand Up @@ -1734,14 +1738,14 @@ def _expand_mapped_task_if_needed(ti: TI) -> Iterable[TI] | None:

# 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()
# (task id, region id) pairs for which _revise_map_indexes_if_mapped was already done
revised_expansion_keys: 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
Expand All @@ -1760,14 +1764,18 @@ def _expand_mapped_task_if_needed(ti: TI) -> Iterable[TI] | None:
# later in this same pass must see the post-expansion count).
dep_context.invalidate_upstream_task_id_counts()
if new_tis is None and schedulable.state in SCHEDULEABLE_STATES:
# It's enough to revise map index once per task id,
# It's enough to revise map index once per task and region,
# 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)
Comment thread
ashb marked this conversation as resolved.
if expansion_key not in revised_expansion_keys:
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_expansion_keys.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
Expand All @@ -1781,10 +1789,9 @@ def _expand_mapped_task_if_needed(ti: TI) -> Iterable[TI] | None:
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

Expand Down Expand Up @@ -2156,7 +2163,12 @@ def _create_task_instances(
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.
Expand All @@ -2181,7 +2193,7 @@ def _revise_map_indexes_if_mapped(
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)
Expand All @@ -2194,7 +2206,7 @@ def _revise_map_indexes_if_mapped(
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)
Expand All @@ -2209,13 +2221,20 @@ def _revise_map_indexes_if_mapped(
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)
Expand Down
Loading
Loading