diff --git a/airflow-core/docs/migrations-ref.rst b/airflow-core/docs/migrations-ref.rst index 0e4ed7a6fae09..009cb049e481d 100644 --- a/airflow-core/docs/migrations-ref.rst +++ b/airflow-core/docs/migrations-ref.rst @@ -39,7 +39,9 @@ Here's the list of all the Database Migrations that are executed via when you ru +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | Revision ID | Revises ID | Airflow Version | Description | +=========================+==================+===================+==============================================================+ -| ``5182d0596ee2`` (head) | ``b6a9c2e7d410`` | ``3.4.0`` | Widen revoked_token.jti to store external-issuer token | +| ``3b7a91c5df20`` (head) | ``5182d0596ee2`` | ``3.4.0`` | Fold task_map into xcom.mapped_length. | ++-------------------------+------------------+-------------------+--------------------------------------------------------------+ +| ``5182d0596ee2`` | ``b6a9c2e7d410`` | ``3.4.0`` | Widen revoked_token.jti to store external-issuer token | | | | | identifiers. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | ``b6a9c2e7d410`` | ``f8c2a1d94e03`` | ``3.4.0`` | Add draining state to DagModel. | diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/xcom.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/xcom.py index 48ed3d11bf199..28304390c8288 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/xcom.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/xcom.py @@ -395,6 +395,8 @@ def update_xcom_entry( run_id=dag_run_id, map_index=patch_body.map_index, serialize=False, + # Not recomputed from the new value: a custom XCom backend stores only a reference. + mapped_length=xcom_entry.mapped_length, session=session, ) except (ValueError, TypeError) as e: diff --git a/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py b/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py index 12336b7386a44..8c6b90b85d59a 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/services/ui/grid.py @@ -27,7 +27,6 @@ from airflow.api_fastapi.common.parameters import state_priority from airflow.api_fastapi.core_api.services.ui.task_group import get_task_group_children_getter -from airflow.models.taskmap import TaskMap from airflow.serialization.definitions.baseoperator import SerializedBaseOperator from airflow.serialization.definitions.mappedoperator import SerializedMappedOperator from airflow.serialization.definitions.taskgroup import SerializedTaskGroup @@ -142,8 +141,8 @@ def _get_aggs_for_node(summary: GridNodeAgg) -> dict[str, Any]: def _find_aggregates( - node: SerializedTaskGroup | SerializedBaseOperator | TaskMap, - parent_node: SerializedTaskGroup | SerializedBaseOperator | TaskMap | None, + node: SerializedTaskGroup | SerializedBaseOperator, + parent_node: SerializedTaskGroup | SerializedBaseOperator | None, ti_details: Mapping[str, GridNodeAgg], group_dict: dict[str | None, SerializedTaskGroup] | None = None, ) -> Iterable[tuple[dict[str, Any], GridNodeAgg]]: diff --git a/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py b/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py index 4dd7f6818fca2..40469b3e56ea4 100644 --- a/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py +++ b/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py @@ -33,8 +33,7 @@ XComSequenceSliceResponse, ) from airflow.api_fastapi.execution_api.security import CurrentTIToken -from airflow.models.taskmap import TaskMap -from airflow.models.xcom import XComModel +from airflow.models.xcom import XCOM_RETURN_KEY, XComModel from airflow.utils.db import get_query_count @@ -397,7 +396,7 @@ def set_xcom( map_index: Annotated[int, Query()] = -1, dag_result: Annotated[bool, Query(description="Whether this XCom is a dag result")] = False, mapped_length: Annotated[ - int | None, Query(description="Number of mapped tasks this value expands into") + int | None, Query(ge=0, description="Number of mapped tasks this value expands into") ] = None, ): """Set an Airflow XCom.""" @@ -415,16 +414,17 @@ def set_xcom( ) if mapped_length is not None: - task_map = TaskMap( - dag_id=dag_id, - task_id=task_id, - run_id=run_id, - map_index=map_index, - length=mapped_length, - keys=None, - ) + # The scheduler only ever reads a length off the return value, so any other key is write-only. + if key != XCOM_RETURN_KEY: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "reason": "invalid_mapped_length_key", + "message": f"mapped_length is only valid for the {XCOM_RETURN_KEY!r} key.", + }, + ) max_map_length = conf.getint("core", "max_map_length", fallback=1024) - if task_map.length > max_map_length: + if mapped_length > max_map_length: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail={ @@ -432,7 +432,6 @@ def set_xcom( "message": "pushed value is too large to map as a downstream's dependency", }, ) - session.merge(task_map) # else: # TODO: Can/should we check if a client _hasn't_ provided this for an upstream of a mapped task? That @@ -449,6 +448,7 @@ def set_xcom( map_index=map_index, serialize=False, dag_result=dag_result, + mapped_length=mapped_length, session=session, ) except ValueError as e: diff --git a/airflow-core/src/airflow/migrations/versions/0136_3_4_0_fold_task_map_into_xcom_mapped_length.py b/airflow-core/src/airflow/migrations/versions/0136_3_4_0_fold_task_map_into_xcom_mapped_length.py new file mode 100644 index 0000000000000..86e1e7ccf4aad --- /dev/null +++ b/airflow-core/src/airflow/migrations/versions/0136_3_4_0_fold_task_map_into_xcom_mapped_length.py @@ -0,0 +1,158 @@ +# +# 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. + +""" +Fold task_map into xcom.mapped_length. + +``task_map.keys`` is dropped rather than migrated, so downgrade restores every +map as the list variant. + +Revision ID: 3b7a91c5df20 +Revises: 5182d0596ee2 +Create Date: 2026-09-10 10:00:00.000000 + +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +from airflow.migrations.db_types import StringID +from airflow.migrations.utils import disable_sqlite_fkeys +from airflow.utils.sqlalchemy import ExtendedJSON + +# revision identifiers, used by Alembic. +revision = "3b7a91c5df20" +down_revision = "5182d0596ee2" +branch_labels = None +depends_on = None +airflow_version = "3.4.0" + +XCOM_RETURN_KEY = "return_value" + + +def _tables(xcom_name: str, task_map_name: str): + xcom = sa.table( + xcom_name, + sa.column("dag_id"), + sa.column("task_id"), + sa.column("run_id"), + sa.column("map_index"), + sa.column("key"), + sa.column("mapped_length"), + ) + task_map = sa.table( + task_map_name, + sa.column("dag_id"), + sa.column("task_id"), + sa.column("run_id"), + sa.column("map_index"), + sa.column("length"), + sa.column("keys"), + ) + join = sa.and_( + task_map.c.dag_id == xcom.c.dag_id, + task_map.c.task_id == xcom.c.task_id, + task_map.c.run_id == xcom.c.run_id, + task_map.c.map_index == xcom.c.map_index, + ) + return xcom, task_map, join + + +def build_backfill_statement(xcom_name: str = "xcom", task_map_name: str = "task_map"): + """ + Copy each task_map length onto the XCom row it describes. + + A correlated subquery rather than UPDATE ... FROM: MySQL has no such form and SQLite only + gained it in 3.33. + """ + xcom, task_map, join = _tables(xcom_name, task_map_name) + return ( + xcom.update() + .where( + xcom.c.key == XCOM_RETURN_KEY, + sa.exists(sa.select(sa.literal(1)).where(join)), + ) + .values(mapped_length=sa.select(task_map.c.length).where(join).scalar_subquery()) + ) + + +def build_restore_statement(xcom_name: str = "xcom", task_map_name: str = "task_map"): + """ + Rebuild task_map rows from the lengths on XCom rows. + + Scoped to one key because task_map's primary key has no key column. + """ + xcom, task_map, _ = _tables(xcom_name, task_map_name) + return task_map.insert().from_select( + ["dag_id", "task_id", "run_id", "map_index", "length", "keys"], + sa.select( + xcom.c.dag_id, + xcom.c.task_id, + xcom.c.run_id, + xcom.c.map_index, + xcom.c.mapped_length, + sa.null(), + ).where(xcom.c.mapped_length.is_not(None), xcom.c.key == XCOM_RETURN_KEY), + ) + + +def upgrade(): + """Fold task_map into xcom.mapped_length.""" + with disable_sqlite_fkeys(op): + with op.batch_alter_table("xcom", schema=None) as batch_op: + batch_op.add_column(sa.Column("mapped_length", sa.Integer(), nullable=True)) + batch_op.create_check_constraint("mapped_length_not_negative", "mapped_length >= 0") + + op.execute(build_backfill_statement()) + op.drop_table("task_map") + + +def downgrade(): + """Restore the task_map table from xcom.mapped_length.""" + with disable_sqlite_fkeys(op): + op.create_table( + "task_map", + sa.Column("dag_id", StringID(length=250), nullable=False), + sa.Column("task_id", StringID(length=250), nullable=False), + sa.Column("run_id", StringID(length=250), nullable=False), + sa.Column("map_index", sa.Integer(), nullable=False), + sa.Column("length", sa.Integer(), nullable=False), + sa.Column("keys", ExtendedJSON(), nullable=True), + sa.CheckConstraint("length >= 0", name="task_map_length_not_negative"), + sa.ForeignKeyConstraint( + ["dag_id", "task_id", "run_id", "map_index"], + [ + "task_instance.dag_id", + "task_instance.task_id", + "task_instance.run_id", + "task_instance.map_index", + ], + name="task_map_task_instance_fkey", + onupdate="CASCADE", + ondelete="CASCADE", + ), + sa.PrimaryKeyConstraint("dag_id", "task_id", "run_id", "map_index", name="task_map_pkey"), + ) + + op.execute(build_restore_statement()) + + with op.batch_alter_table("xcom", schema=None) as batch_op: + batch_op.drop_constraint("mapped_length_not_negative", type_="check") + batch_op.drop_column("mapped_length") diff --git a/airflow-core/src/airflow/models/dagrun.py b/airflow-core/src/airflow/models/dagrun.py index 223fd3d05ee60..2d7c8de4de7d0 100644 --- a/airflow-core/src/airflow/models/dagrun.py +++ b/airflow-core/src/airflow/models/dagrun.py @@ -91,7 +91,6 @@ from airflow.models.taskinstance import TaskInstance as TI, _add_and_prime_mapped_ti, clear_task_instances from airflow.models.taskinstancehistory import TaskInstanceHistory as TIH from airflow.models.tasklog import LogTemplate -from airflow.models.taskmap import TaskMap from airflow.serialization.definitions.deadline import SerializedReferenceModels from airflow.serialization.definitions.notset import NOTSET, ArgNotSet, is_arg_set from airflow.ti_deps.dep_context import DepContext @@ -1729,7 +1728,7 @@ def _expand_mapped_task_if_needed(ti: TI) -> Iterable[TI] | None: # the db references. ti.clear_db_references(session=session) try: - expanded_tis, _ = TaskMap.expand_mapped_task(ti.task, self.run_id, session=session) + expanded_tis, _ = ti.expand_mapped_task(session=session) except NotMapped: # Not a mapped task, nothing needed. return None if expanded_tis: diff --git a/airflow-core/src/airflow/models/taskinstance.py b/airflow-core/src/airflow/models/taskinstance.py index d9c3f8cab9595..327c17028b1a1 100644 --- a/airflow-core/src/airflow/models/taskinstance.py +++ b/airflow-core/src/airflow/models/taskinstance.py @@ -91,7 +91,6 @@ from airflow.models.hitl import HITLDetail # noqa: F401 from airflow.models.log import Log from airflow.models.taskinstancekey import TaskInstanceKey -from airflow.models.taskmap import TaskMap from airflow.models.taskreschedule import TaskReschedule from airflow.models.xcom import XCOM_RETURN_KEY, LazyXComSelectSequence, XComModel from airflow.serialization.enums import stringify_encoding_keys @@ -100,12 +99,13 @@ from airflow.ti_deps.dep_context import DepContext from airflow.ti_deps.dependencies_deps import REQUEUEABLE_DEPS, RUNNING_DEPS from airflow.ti_deps.deps.ready_to_reschedule import ReadyToRescheduleDep +from airflow.utils.db import exists_query from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.net import get_hostname from airflow.utils.platform import getuser from airflow.utils.retries import run_with_db_retries from airflow.utils.session import NEW_SESSION, create_session, provide_session -from airflow.utils.sqlalchemy import ExecutorConfigType, ExtendedJSON, UtcDateTime +from airflow.utils.sqlalchemy import ExecutorConfigType, ExtendedJSON, UtcDateTime, with_row_locks from airflow.utils.state import DagRunState, State, TaskInstanceState TR = TaskReschedule @@ -2326,7 +2326,6 @@ def clear_db_references(self, session: Session): tables: list[type[TaskInstanceDependencies]] = [ XComModel, RenderedTaskInstanceFields, - TaskMap, ] tables_by_id: list[type[Base]] = [TaskInstanceNote, TaskReschedule] for table in tables: @@ -2341,6 +2340,167 @@ def clear_db_references(self, session: Session): for table in tables_by_id: session.execute(delete(table).where(table.ti_id == self.id)) + def expand_mapped_task(self, *, session: Session) -> tuple[Sequence[TaskInstance], int]: + """ + Create the mapped task instances for mapped task. + + :raise NotMapped: If this task does not need expansion. + :return: The newly created mapped task instances (if any) in ascending + order by map index, and the maximum map index value. + """ + from airflow.models.expandinput import NotFullyPopulated + from airflow.serialization.definitions.baseoperator import SerializedBaseOperator + from airflow.serialization.definitions.mappedoperator import ( + SerializedMappedOperator, + get_mapped_ti_count, + ) + + task = self.task + run_id = self.run_id + + if not isinstance(task, (SerializedMappedOperator, SerializedBaseOperator)): + raise RuntimeError( + f"cannot expand unrecognized operator type {type(task).__module__}.{type(task).__name__}" + ) + + try: + total_length: int | None = get_mapped_ti_count(task, run_id, session=session) + except NotFullyPopulated as e: + if not task.dag or not task.dag.partial: + task.log.error( + "Cannot expand %r for run %s; missing upstream values: %s", + task, + run_id, + sorted(e.missing), + ) + total_length = None + + state: str | None = None + unmapped_ti: TaskInstance | None = session.scalars( + select(TaskInstance).where( + TaskInstance.dag_id == task.dag_id, + TaskInstance.task_id == task.task_id, + TaskInstance.run_id == run_id, + TaskInstance.map_index == -1, + or_(TaskInstance.state.in_(State.unfinished), TaskInstance.state.is_(None)), + ) + ).one_or_none() + + all_expanded_tis: list[TaskInstance] = [] + + if unmapped_ti: + if TYPE_CHECKING: + assert task.dag is None + + # The unmapped task instance still exists and is unfinished, i.e. we + # haven't tried to run it before. + if total_length is None: + # If the DAG is partial, it's likely that the upstream tasks + # are not done yet, so the task can't fail yet. + if not task.dag or not task.dag.partial: + unmapped_ti.state = TaskInstanceState.UPSTREAM_FAILED + elif total_length < 1: + # If the upstream maps this to a zero-length value, simply mark + # the unmapped task instance as SKIPPED (if needed). + task.log.info( + "Marking %s as SKIPPED since the map has %d values to expand", + unmapped_ti, + total_length, + ) + unmapped_ti.state = TaskInstanceState.SKIPPED + else: + dr = unmapped_ti.dag_run + zero_index_ti_exists = exists_query( + TaskInstance.dag_id == task.dag_id, + TaskInstance.task_id == task.task_id, + TaskInstance.run_id == run_id, + TaskInstance.map_index == 0, + session=session, + ) + if not zero_index_ti_exists: + # Otherwise convert this into the first mapped index, and create + # TaskInstance for other indexes. + unmapped_ti.map_index = 0 + task.log.debug("Updated in place to become %s", unmapped_ti) + all_expanded_tis.append(unmapped_ti) + # execute hook for task instance map index 0 + task_instance_mutation_hook(unmapped_ti, dag_run=dr) + session.flush() + else: + task.log.debug("Deleting the original task instance: %s", unmapped_ti) + session.delete(unmapped_ti) + state = unmapped_ti.state + + if total_length is None or total_length < 1: + # Nothing to fixup. + indexes_to_map: Iterable[int] = () + else: + # Only create "missing" ones. + current_max_mapping = ( + session.scalar( + select(func.max(TaskInstance.map_index)).where( + TaskInstance.dag_id == task.dag_id, + TaskInstance.task_id == task.task_id, + TaskInstance.run_id == run_id, + ) + ) + or 0 + ) + indexes_to_map = range(current_max_mapping + 1, total_length) + + if unmapped_ti: + dag_version_id = unmapped_ti.dag_version_id + elif dag_version := DagVersion.get_latest_version(task.dag_id, session=session): + dag_version_id = dag_version.id + else: + dag_version_id = None + + if not unmapped_ti: + from airflow.models import DagRun + + dr = session.scalar( + select(DagRun).where( + DagRun.dag_id == task.dag_id, + DagRun.run_id == run_id, + ) + ) + + new_tis: list[TaskInstance] = [] + for index in indexes_to_map: + ti = TaskInstance( + task, + run_id=run_id, + map_index=index, + state=state, + dag_version_id=dag_version_id, + ) + task.log.debug("Expanding TIs upserted %s", ti) + _add_and_prime_mapped_ti( + ti, task, dr, session=session, context_carrier=new_task_run_carrier(dr.context_carrier) + ) + new_tis.append(ti) + if new_tis: + session.flush() + all_expanded_tis.extend(new_tis) + + # Coerce the None case to 0 -- these two are almost treated identically, + # except the unmapped ti (if exists) is marked to different states. + total_expanded_ti_count = total_length or 0 + + # Any (old) task instances with inapplicable indexes (>= the total + # number we need) are set to "REMOVED". + query = select(TaskInstance).where( + TaskInstance.dag_id == task.dag_id, + TaskInstance.task_id == task.task_id, + TaskInstance.run_id == run_id, + TaskInstance.map_index >= total_expanded_ti_count, + ) + to_update = session.scalars(with_row_locks(query, of=TaskInstance, session=session, skip_locked=True)) + for ti in to_update: + ti.state = TaskInstanceState.REMOVED + session.flush() + return all_expanded_tis, total_expanded_ti_count - 1 + @classmethod def duration_expression_update( cls, end_date: datetime, query: Update, bind: Engine | SAConnection diff --git a/airflow-core/src/airflow/models/taskmap.py b/airflow-core/src/airflow/models/taskmap.py deleted file mode 100644 index 96b3c0831cbcb..0000000000000 --- a/airflow-core/src/airflow/models/taskmap.py +++ /dev/null @@ -1,293 +0,0 @@ -# -# 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. -"""Table to store information about mapped task instances (AIP-42).""" - -from __future__ import annotations - -import collections.abc -import enum -from collections.abc import Collection, Iterable, Sequence -from typing import TYPE_CHECKING, Any - -from opentelemetry import trace -from sqlalchemy import CheckConstraint, ForeignKeyConstraint, Integer, String, func, or_, select -from sqlalchemy.orm import Mapped, mapped_column - -from airflow._shared.observability.traces import new_task_run_carrier -from airflow.models.base import COLLATION_ARGS, ID_LEN, TaskInstanceDependencies -from airflow.models.dag_version import DagVersion -from airflow.utils.db import exists_query -from airflow.utils.sqlalchemy import ExtendedJSON, with_row_locks -from airflow.utils.state import State, TaskInstanceState - -if TYPE_CHECKING: - from sqlalchemy.orm import Session - - from airflow.models.taskinstance import TaskInstance - from airflow.serialization.definitions.mappedoperator import Operator -tracer = trace.get_tracer(__name__) - - -class TaskMapVariant(enum.Enum): - """ - Task map variant. - - Possible values are **dict** (for a key-value mapping) and **list** (for an - ordered value sequence). - """ - - DICT = "dict" - LIST = "list" - - -class TaskMap(TaskInstanceDependencies): - """ - Model to track dynamic task-mapping information. - - This is currently only populated by an upstream TaskInstance pushing an - XCom that's pulled by a downstream for mapping purposes. - """ - - __tablename__ = "task_map" - - # Link to upstream TaskInstance creating this dynamic mapping information. - dag_id: Mapped[str] = mapped_column(String(ID_LEN, **COLLATION_ARGS), primary_key=True) - task_id: Mapped[str] = mapped_column(String(ID_LEN, **COLLATION_ARGS), primary_key=True) - run_id: Mapped[str] = mapped_column(String(ID_LEN, **COLLATION_ARGS), primary_key=True) - map_index: Mapped[int] = mapped_column(Integer, primary_key=True) - - length: Mapped[int] = mapped_column(Integer, nullable=False) - keys: Mapped[list | None] = mapped_column(ExtendedJSON, nullable=True) - - __table_args__ = ( - CheckConstraint(length >= 0, name="task_map_length_not_negative"), - ForeignKeyConstraint( - [dag_id, task_id, run_id, map_index], - [ - "task_instance.dag_id", - "task_instance.task_id", - "task_instance.run_id", - "task_instance.map_index", - ], - name="task_map_task_instance_fkey", - ondelete="CASCADE", - onupdate="CASCADE", - ), - ) - - def __init__( - self, - dag_id: str, - task_id: str, - run_id: str, - map_index: int, - length: int, - keys: list[Any] | None, - ) -> None: - self.dag_id = dag_id - self.task_id = task_id - self.run_id = run_id - self.map_index = map_index - self.length = length - self.keys = keys - - @classmethod - def from_task_instance_xcom(cls, ti: TaskInstance, value: Collection) -> TaskMap: - if ti.run_id is None: - raise ValueError("cannot record task map for unrun task instance") - return cls( - dag_id=ti.dag_id, - task_id=ti.task_id, - run_id=ti.run_id, - map_index=ti.map_index, - length=len(value), - keys=(list(value) if isinstance(value, collections.abc.Mapping) else None), - ) - - @property - def variant(self) -> TaskMapVariant: - if self.keys is None: - return TaskMapVariant.LIST - return TaskMapVariant.DICT - - @classmethod - def expand_mapped_task( - cls, - task: Operator, - run_id: str, - *, - session: Session, - ) -> tuple[Sequence[TaskInstance], int]: - """ - Create the mapped task instances for mapped task. - - :raise NotMapped: If this task does not need expansion. - :return: The newly created mapped task instances (if any) in ascending - order by map index, and the maximum map index value. - """ - from airflow.models.expandinput import NotFullyPopulated - from airflow.models.taskinstance import TaskInstance, _add_and_prime_mapped_ti - from airflow.serialization.definitions.baseoperator import SerializedBaseOperator - from airflow.serialization.definitions.mappedoperator import ( - SerializedMappedOperator, - get_mapped_ti_count, - ) - from airflow.settings import task_instance_mutation_hook - - if not isinstance(task, (SerializedMappedOperator, SerializedBaseOperator)): - raise RuntimeError( - f"cannot expand unrecognized operator type {type(task).__module__}.{type(task).__name__}" - ) - - try: - total_length: int | None = get_mapped_ti_count(task, run_id, session=session) - except NotFullyPopulated as e: - if not task.dag or not task.dag.partial: - task.log.error( - "Cannot expand %r for run %s; missing upstream values: %s", - task, - run_id, - sorted(e.missing), - ) - total_length = None - - state: str | None = None - unmapped_ti: TaskInstance | None = session.scalars( - select(TaskInstance).where( - TaskInstance.dag_id == task.dag_id, - TaskInstance.task_id == task.task_id, - TaskInstance.run_id == run_id, - TaskInstance.map_index == -1, - or_(TaskInstance.state.in_(State.unfinished), TaskInstance.state.is_(None)), - ) - ).one_or_none() - - all_expanded_tis: list[TaskInstance] = [] - - if unmapped_ti: - if TYPE_CHECKING: - assert task.dag is None - - # The unmapped task instance still exists and is unfinished, i.e. we - # haven't tried to run it before. - if total_length is None: - # If the DAG is partial, it's likely that the upstream tasks - # are not done yet, so the task can't fail yet. - if not task.dag or not task.dag.partial: - unmapped_ti.state = TaskInstanceState.UPSTREAM_FAILED - elif total_length < 1: - # If the upstream maps this to a zero-length value, simply mark - # the unmapped task instance as SKIPPED (if needed). - task.log.info( - "Marking %s as SKIPPED since the map has %d values to expand", - unmapped_ti, - total_length, - ) - unmapped_ti.state = TaskInstanceState.SKIPPED - else: - dr = unmapped_ti.dag_run - zero_index_ti_exists = exists_query( - TaskInstance.dag_id == task.dag_id, - TaskInstance.task_id == task.task_id, - TaskInstance.run_id == run_id, - TaskInstance.map_index == 0, - session=session, - ) - if not zero_index_ti_exists: - # Otherwise convert this into the first mapped index, and create - # TaskInstance for other indexes. - unmapped_ti.map_index = 0 - task.log.debug("Updated in place to become %s", unmapped_ti) - all_expanded_tis.append(unmapped_ti) - # execute hook for task instance map index 0 - task_instance_mutation_hook(unmapped_ti, dag_run=dr) - session.flush() - else: - task.log.debug("Deleting the original task instance: %s", unmapped_ti) - session.delete(unmapped_ti) - state = unmapped_ti.state - - if total_length is None or total_length < 1: - # Nothing to fixup. - indexes_to_map: Iterable[int] = () - else: - # Only create "missing" ones. - current_max_mapping = ( - session.scalar( - select(func.max(TaskInstance.map_index)).where( - TaskInstance.dag_id == task.dag_id, - TaskInstance.task_id == task.task_id, - TaskInstance.run_id == run_id, - ) - ) - or 0 - ) - indexes_to_map = range(current_max_mapping + 1, total_length) - - if unmapped_ti: - dag_version_id = unmapped_ti.dag_version_id - elif dag_version := DagVersion.get_latest_version(task.dag_id, session=session): - dag_version_id = dag_version.id - else: - dag_version_id = None - - if not unmapped_ti: - from airflow.models import DagRun - - dr = session.scalar( - select(DagRun).where( - DagRun.dag_id == task.dag_id, - DagRun.run_id == run_id, - ) - ) - - new_tis: list[TaskInstance] = [] - for index in indexes_to_map: - ti = TaskInstance( - task, - run_id=run_id, - map_index=index, - state=state, - dag_version_id=dag_version_id, - ) - task.log.debug("Expanding TIs upserted %s", ti) - _add_and_prime_mapped_ti( - ti, task, dr, session=session, context_carrier=new_task_run_carrier(dr.context_carrier) - ) - new_tis.append(ti) - if new_tis: - session.flush() - all_expanded_tis.extend(new_tis) - - # Coerce the None case to 0 -- these two are almost treated identically, - # except the unmapped ti (if exists) is marked to different states. - total_expanded_ti_count = total_length or 0 - - # Any (old) task instances with inapplicable indexes (>= the total - # number we need) are set to "REMOVED". - query = select(TaskInstance).where( - TaskInstance.dag_id == task.dag_id, - TaskInstance.task_id == task.task_id, - TaskInstance.run_id == run_id, - TaskInstance.map_index >= total_expanded_ti_count, - ) - to_update = session.scalars(with_row_locks(query, of=TaskInstance, session=session, skip_locked=True)) - for ti in to_update: - ti.state = TaskInstanceState.REMOVED - session.flush() - return all_expanded_tis, total_expanded_ti_count - 1 diff --git a/airflow-core/src/airflow/models/xcom.py b/airflow-core/src/airflow/models/xcom.py index 5b27f244bf517..58e5c429f66b4 100644 --- a/airflow-core/src/airflow/models/xcom.py +++ b/airflow-core/src/airflow/models/xcom.py @@ -26,6 +26,7 @@ from sqlalchemy import ( JSON, Boolean, + CheckConstraint, ForeignKeyConstraint, Index, Integer, @@ -76,12 +77,16 @@ class XComModel(TaskInstanceDependencies): value: Mapped[Any] = mapped_column(JSON().with_variant(postgresql.JSONB, "postgresql"), nullable=True) timestamp: Mapped[datetime] = mapped_column(UtcDateTime, default=timezone.utcnow, nullable=False) + # NULL unless the value can expand a downstream mapped task (AIP-42). + mapped_length: Mapped[int | None] = mapped_column(Integer, nullable=True) + __table_args__ = ( # Ideally we should create a unique index over (key, dag_id, task_id, run_id), # but it goes over MySQL's index length limit. So we instead index 'key' # separately, and enforce uniqueness with DagRun.id instead. Index("idx_xcom_key", key), Index("idx_xcom_task_instance", dag_id, task_id, run_id, map_index), + CheckConstraint(mapped_length >= 0, name="mapped_length_not_negative"), PrimaryKeyConstraint("dag_run_id", "task_id", "map_index", "key", name="xcom_pkey"), ForeignKeyConstraint( [dag_id, task_id, run_id, map_index], @@ -166,6 +171,7 @@ def set( map_index: int = -1, serialize: bool = True, dag_result: bool = False, + mapped_length: int | None = None, session: Session = NEW_SESSION, ) -> None: """ @@ -179,6 +185,8 @@ def set( :param map_index: Optional map index to assign XCom for a mapped task. :param serialize: Optional parameter to specify if value should be serialized or not. The default is ``True``. + :param mapped_length: Length of the value, if it can be used to expand a + downstream mapped task. :param session: Database session. If not given, a new session will be created for this function. """ @@ -245,6 +253,7 @@ def set( dag_id=dag_id, map_index=map_index, dag_result=dag_result, + mapped_length=mapped_length, ) session.add(new) session.flush() diff --git a/airflow-core/src/airflow/serialization/definitions/xcom_arg.py b/airflow-core/src/airflow/serialization/definitions/xcom_arg.py index ebca6f2fa5193..cd544ae5c7bb8 100644 --- a/airflow-core/src/airflow/serialization/definitions/xcom_arg.py +++ b/airflow-core/src/airflow/serialization/definitions/xcom_arg.py @@ -154,7 +154,6 @@ def get_task_map_length(xcom_arg: SchedulerXComArg, run_id: str, *, session: Ses @get_task_map_length.register def _(xcom_arg: SchedulerPlainXComArg, run_id: str, *, session: Session) -> int | None: from airflow.models.taskinstance import TaskInstance - from airflow.models.taskmap import TaskMap from airflow.models.xcom import XComModel from airflow.serialization.definitions.mappedoperator import is_mapped @@ -177,21 +176,26 @@ 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. - query = select(func.count(XComModel.map_index)).where( + return session.scalar( + select(func.count(XComModel.map_index)).where( + XComModel.dag_id == dag_id, + XComModel.run_id == run_id, + XComModel.task_id == task_id, + XComModel.map_index >= 0, + XComModel.key == XCOM_RETURN_KEY, + ) + ) + + # Not xcom_arg.key: the SDK records the length of the whole return value, never per key. + return session.scalar( + select(XComModel.mapped_length).where( XComModel.dag_id == dag_id, XComModel.run_id == run_id, XComModel.task_id == task_id, - XComModel.map_index >= 0, + XComModel.map_index == -1, XComModel.key == XCOM_RETURN_KEY, ) - else: - query = select(TaskMap.length).where( - TaskMap.dag_id == dag_id, - TaskMap.run_id == run_id, - TaskMap.task_id == task_id, - TaskMap.map_index < 0, - ) - return session.scalar(query) + ) @get_task_map_length.register diff --git a/airflow-core/src/airflow/utils/db.py b/airflow-core/src/airflow/utils/db.py index 23945ebecd4f8..a116050d1ba7d 100644 --- a/airflow-core/src/airflow/utils/db.py +++ b/airflow-core/src/airflow/utils/db.py @@ -117,7 +117,7 @@ class MappedClassProtocol(Protocol): "3.1.8": "509b94a1042d", "3.2.0": "1d6611b6ab7c", "3.3.0": "d2f4e1b3c5a7", - "3.4.0": "5182d0596ee2", + "3.4.0": "3b7a91c5df20", } # Prefix used to identify tables holding data moved during migration. diff --git a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py index 4e51476494e10..53f5bc85aa3e0 100644 --- a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py +++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py @@ -45,7 +45,6 @@ from airflow.models.task_state_store import TaskStateStoreModel from airflow.models.taskinstance import uuid7 from airflow.models.taskinstancehistory import TaskInstanceHistory -from airflow.models.taskmap import TaskMap from airflow.models.team import Team from airflow.models.trigger import Trigger from airflow.sdk import BaseOperator, TaskGroup @@ -63,6 +62,7 @@ clear_rendered_ti_fields, ) from tests_common.test_utils.logs import check_last_log +from tests_common.test_utils.mapping import expand_mapped_task_instances, push_mapped_length from tests_common.test_utils.mock_operators import MockOperator from tests_common.test_utils.taskinstance import create_task_instance from tests_common.test_utils.team import attach_dag_to_team @@ -733,15 +733,8 @@ def create_dag_runs_with_mapped_tasks(self, dag_maker, session, dags=None): data_interval=(DEFAULT_DATETIME_1, DEFAULT_DATETIME_2), ) dag_version = DagVersion.get_latest_version(dag_id) - session.add( - TaskMap( - dag_id=dr.dag_id, - task_id=task1.task_id, - run_id=dr.run_id, - map_index=-1, - length=count, - keys=None, - ) + push_mapped_length( + dr.get_task_instance(task1.task_id, session=session), list(range(count)), session=session ) if count: @@ -773,7 +766,7 @@ def create_dag_runs_with_mapped_tasks(self, dag_maker, session, dags=None): sync_bag_to_db(dagbag, "dags-folder", None) session.flush() - TaskMap.expand_mapped_task(sdag.task_dict[mapped.task_id], dr.run_id, session=session) + expand_mapped_task_instances(sdag.task_dict[mapped.task_id], dr.run_id, session=session) @pytest.fixture def one_task_with_mapped_tis(self, dag_maker, session): diff --git a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_xcom.py b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_xcom.py index d3d59ce9f48ae..93e7f40cfec5c 100644 --- a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_xcom.py +++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_xcom.py @@ -22,7 +22,7 @@ from unittest import mock import pytest -from sqlalchemy import update +from sqlalchemy import select, update from sqlalchemy.orm import Session from airflow._shared.timezones import timezone @@ -32,7 +32,7 @@ from airflow.models.dagbundle import DagBundleModel from airflow.models.dagrun import DagRun from airflow.models.team import Team -from airflow.models.xcom import XComModel +from airflow.models.xcom import XCOM_RETURN_KEY, XComModel from airflow.providers.standard.operators.empty import EmptyOperator from airflow.sdk import DAG, AssetAlias from airflow.sdk.bases.xcom import BaseXCom @@ -1086,3 +1086,36 @@ def test_patch_xcom_preserves_int_type(self, test_client, session): assert data["value"] == patch_value assert isinstance(data["value"], int), f"Expected int type but got {type(data['value'])}" check_last_log(session, dag_id=TEST_DAG_ID, event="update_xcom_entry", logical_date=None) + + def test_patch_xcom_preserves_mapped_length(self, test_client, session): + """set() replaces the row, so an edit must not drop the recorded expansion length.""" + key = XCOM_RETURN_KEY + XComModel.set( + key=key, + value=[1, 2, 3], + dag_id=TEST_DAG_ID, + task_id=TEST_TASK_ID, + run_id=run_id, + mapped_length=3, + session=session, + ) + session.commit() + + response = test_client.patch( + f"/dags/{TEST_DAG_ID}/dagRuns/{run_id}/taskInstances/{TEST_TASK_ID}/xcomEntries/{key}", + json={"value": [9, 9, 9]}, + ) + + assert response.status_code == 200 + assert response.json()["value"] == [9, 9, 9] + assert ( + session.scalar( + select(XComModel.mapped_length).where( + XComModel.dag_id == TEST_DAG_ID, + XComModel.task_id == TEST_TASK_ID, + XComModel.run_id == run_id, + XComModel.key == key, + ) + ) + == 3 + ) diff --git a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py index 8ecf7b1b5dd17..534e203490511 100644 --- a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py +++ b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py @@ -608,7 +608,7 @@ def expandable_task_group(param: str) -> None: def test_dynamic_task_mapping_with_xcom(self, client: Client, dag_maker: DagMaker, session: Session): """Test that dynamic task mapping works correctly with XCom values.""" - from airflow.models.taskmap import TaskMap + from tests_common.test_utils.mapping import push_mapped_length with dag_maker(session=session, serialized=True): @@ -634,10 +634,10 @@ def task_3(): decision = dr.task_instance_scheduling_decisions(session=session) - # Simulate task_1 execution to produce TaskMap. + # Simulate task_1 execution to produce the mapped length. (ti_1,) = decision.schedulable_tis ti_1.state = TaskInstanceState.SUCCESS - session.add(TaskMap.from_task_instance_xcom(ti_1, [0, 1])) + push_mapped_length(ti_1, [0, 1], session=session) session.flush() # Now task_2 in mapped tagk group is expanded. diff --git a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_xcoms.py b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_xcoms.py index 64ddba31fc0e2..559bfe11aff8d 100644 --- a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_xcoms.py +++ b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_xcoms.py @@ -31,8 +31,7 @@ from airflow.api_fastapi.execution_api.datamodels.xcom import XComResponse from airflow.api_fastapi.execution_api.security import require_auth from airflow.models.dagrun import DagRun -from airflow.models.taskmap import TaskMap -from airflow.models.xcom import XComModel +from airflow.models.xcom import XCOM_RETURN_KEY, XComModel from airflow.providers.standard.operators.empty import EmptyOperator from airflow.sdk.serde import deserialize, serialize from airflow.utils.session import create_session @@ -449,10 +448,7 @@ def test_xcom_set(self, client, create_task_instance, session, value, expected_v ) ).first() assert xcom.value == expected_value - task_map = session.scalars( - select(TaskMap).where(TaskMap.task_id == ti.task_id, TaskMap.dag_id == ti.dag_id) - ).one_or_none() - assert task_map is None, "Should not be mapped" + assert xcom.mapped_length is None, "Should not be mapped" @pytest.mark.parametrize( ("orig_value", "ser_value", "deser_value"), @@ -525,7 +521,7 @@ def test_xcom_set_mapped(self, client, create_task_instance, session): value = serialize("value1") response = client.post( - f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/xcom_1", + f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/{XCOM_RETURN_KEY}", params={"map_index": -1, "mapped_length": 3}, json=value, ) @@ -537,20 +533,41 @@ def test_xcom_set_mapped(self, client, create_task_instance, session): select(XComModel).where( XComModel.task_id == ti.task_id, XComModel.dag_id == ti.dag_id, - XComModel.key == "xcom_1", + XComModel.key == XCOM_RETURN_KEY, XComModel.map_index == -1, ) ).first() assert xcom.value == "value1" - task_map = session.scalars( - select(TaskMap).where(TaskMap.task_id == ti.task_id, TaskMap.dag_id == ti.dag_id) - ).one_or_none() - assert task_map is not None, "Should be mapped" - assert task_map.dag_id == "dag" - assert task_map.run_id == "test" - assert task_map.task_id == "op1" - assert task_map.map_index == -1 - assert task_map.length == 3 + assert xcom.dag_id == "dag" + assert xcom.run_id == "test" + assert xcom.task_id == "op1" + assert xcom.map_index == -1 + assert xcom.mapped_length == 3 + + def test_xcom_set_mapped_rejects_non_return_value_key(self, client, create_task_instance, session): + """Only the return value expands a downstream, so a length under any other key is refused.""" + ti = create_task_instance() + session.commit() + + response = client.post( + f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/xcom_1", + params={"map_index": -1, "mapped_length": 3}, + json=serialize("value1"), + ) + + assert response.status_code == 400 + assert response.json()["detail"]["reason"] == "invalid_mapped_length_key" + + assert ( + session.scalars( + select(XComModel).where( + XComModel.task_id == ti.task_id, + XComModel.dag_id == ti.dag_id, + XComModel.key == "xcom_1", + ) + ).one_or_none() + is None + ) @pytest.mark.parametrize( ("length", "expected_status"), @@ -571,17 +588,23 @@ def test_xcom_set_downstream_of_mapped( session.commit() response = client.post( - f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/xcom_1", + f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/{XCOM_RETURN_KEY}", json='"valid json"', params={"mapped_length": length}, ) assert response.status_code == expected_status + xcom = session.scalars( + select(XComModel).where( + XComModel.task_id == ti.task_id, + XComModel.dag_id == ti.dag_id, + XComModel.key == XCOM_RETURN_KEY, + ) + ).one_or_none() if expected_status < 400: - task_map = session.scalars( - select(TaskMap).where(TaskMap.task_id == ti.task_id, TaskMap.dag_id == ti.dag_id) - ).one_or_none() - assert task_map.length == length + assert xcom.mapped_length == length + else: + assert xcom is None, "Nothing should be written when the length is rejected" @pytest.mark.usefixtures("access_denied") def test_xcom_access_denied(self, client, caplog): diff --git a/airflow-core/tests/unit/migrations/test_0136_fold_task_map_into_xcom.py b/airflow-core/tests/unit/migrations/test_0136_fold_task_map_into_xcom.py new file mode 100644 index 0000000000000..1b2e575a502b2 --- /dev/null +++ b/airflow-core/tests/unit/migrations/test_0136_fold_task_map_into_xcom.py @@ -0,0 +1,184 @@ +# +# 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. + +""" +Tests for migration 0136 (3b7a91c5df20), which folds task_map into xcom.mapped_length. + +An in-flight DagRun's expansion length either survives the backfill or quietly disappears, +so the real statements run against isolated tables on whichever backend the suite is using. +""" + +from __future__ import annotations + +import importlib.util +from pathlib import Path + +import pytest +import sqlalchemy as sa + +from airflow import settings + +from tests_common.test_utils.paths import AIRFLOW_CORE_SOURCES_PATH + +pytestmark = pytest.mark.db_test + +# Migration filenames start with a digit so they cannot be imported via the normal import +# system; load the module by file path instead. +_MIGRATION_PATH = ( + Path(AIRFLOW_CORE_SOURCES_PATH) + / "airflow/migrations/versions/0136_3_4_0_fold_task_map_into_xcom_mapped_length.py" +) +_spec = importlib.util.spec_from_file_location("migration_0134", _MIGRATION_PATH) +_migration = importlib.util.module_from_spec(_spec) # type: ignore[arg-type] +_spec.loader.exec_module(_migration) # type: ignore[union-attr] + +# Isolated because the live xcom table has FK and NOT NULL columns, and task_map is dropped. +_XCOM = "_test_xcom_0134" +_TASK_MAP = "_test_task_map_0134" + +_RETURN_VALUE = "return_value" + +_metadata = sa.MetaData() +_xcom = sa.Table( + _XCOM, + _metadata, + sa.Column("dag_id", sa.String(250)), + sa.Column("task_id", sa.String(250)), + sa.Column("run_id", sa.String(250)), + sa.Column("map_index", sa.Integer), + # Declared rather than raw DDL so SQLAlchemy quotes it: reserved on MySQL. + sa.Column("key", sa.String(512)), + sa.Column("mapped_length", sa.Integer), +) +_task_map = sa.Table( + _TASK_MAP, + _metadata, + sa.Column("dag_id", sa.String(250), nullable=False), + sa.Column("task_id", sa.String(250), nullable=False), + sa.Column("run_id", sa.String(250), nullable=False), + sa.Column("map_index", sa.Integer, nullable=False), + sa.Column("length", sa.Integer, nullable=False), + sa.Column("keys", sa.String(512)), + # The key-less PK the downgrade recreates, so an unscoped restore collides here. + sa.PrimaryKeyConstraint("dag_id", "task_id", "run_id", "map_index"), +) + +_FEEDS_MAPPED = ("d", "feeds_mapped", "r", -1) +_PLAIN = ("d", "plain", "r", -1) + + +def _xcom_row(coords, key, mapped_length=None): + dag_id, task_id, run_id, map_index = coords + return { + "dag_id": dag_id, + "task_id": task_id, + "run_id": run_id, + "map_index": map_index, + "key": key, + "mapped_length": mapped_length, + } + + +def _task_map_row(coords, length): + dag_id, task_id, run_id, map_index = coords + return { + "dag_id": dag_id, + "task_id": task_id, + "run_id": run_id, + "map_index": map_index, + "length": length, + "keys": None, + } + + +@pytest.fixture +def conn(): + _metadata.drop_all(settings.engine) + _metadata.create_all(settings.engine) + try: + with settings.engine.begin() as connection: + yield connection + finally: + _metadata.drop_all(settings.engine) + + +def _lengths(conn) -> dict[tuple[str, str], int | None]: + rows = conn.execute(sa.select(_xcom.c.task_id, _xcom.c.key, _xcom.c.mapped_length)).all() + return {(r.task_id, r.key): r.mapped_length for r in rows} + + +def test_backfill_copies_the_length_onto_the_return_value_row(conn): + conn.execute( + _xcom.insert(), + [ + _xcom_row(_FEEDS_MAPPED, _RETURN_VALUE), + _xcom_row(_FEEDS_MAPPED, "side_output"), + _xcom_row(_PLAIN, _RETURN_VALUE), + ], + ) + conn.execute(_task_map.insert(), [_task_map_row(_FEEDS_MAPPED, 3)]) + + conn.execute(_migration.build_backfill_statement(_XCOM, _TASK_MAP)) + + assert _lengths(conn) == { + ("feeds_mapped", _RETURN_VALUE): 3, + ("feeds_mapped", "side_output"): None, + ("plain", _RETURN_VALUE): None, + } + + +def test_backfill_is_idempotent(conn): + conn.execute(_xcom.insert(), [_xcom_row(_FEEDS_MAPPED, _RETURN_VALUE)]) + conn.execute(_task_map.insert(), [_task_map_row(_FEEDS_MAPPED, 3)]) + + conn.execute(_migration.build_backfill_statement(_XCOM, _TASK_MAP)) + conn.execute(_migration.build_backfill_statement(_XCOM, _TASK_MAP)) + + assert _lengths(conn) == {("feeds_mapped", _RETURN_VALUE): 3} + + +def test_restore_rebuilds_task_map_from_the_return_value_length(conn): + conn.execute( + _xcom.insert(), + [ + _xcom_row(_FEEDS_MAPPED, _RETURN_VALUE, mapped_length=3), + _xcom_row(_PLAIN, _RETURN_VALUE), + ], + ) + + conn.execute(_migration.build_restore_statement(_XCOM, _TASK_MAP)) + + # Subscripted because ``.c.keys`` would resolve to ColumnCollection.keys, the method. + assert conn.execute(sa.select(_task_map.c.task_id, _task_map.c.length, _task_map.c["keys"])).all() == [ + ("feeds_mapped", 3, None) + ] + + +def test_restore_ignores_a_length_recorded_under_another_key(conn): + """Restoring both keys would collide on task_map's key-less primary key.""" + conn.execute( + _xcom.insert(), + [ + _xcom_row(_FEEDS_MAPPED, _RETURN_VALUE, mapped_length=3), + _xcom_row(_FEEDS_MAPPED, "side_output", mapped_length=9), + ], + ) + + conn.execute(_migration.build_restore_statement(_XCOM, _TASK_MAP)) + + assert conn.execute(sa.select(_task_map.c.task_id, _task_map.c.length)).all() == [("feeds_mapped", 3)] diff --git a/airflow-core/tests/unit/models/test_dagrun.py b/airflow-core/tests/unit/models/test_dagrun.py index 850952ca182fe..e53cb7a0318e6 100644 --- a/airflow-core/tests/unit/models/test_dagrun.py +++ b/airflow-core/tests/unit/models/test_dagrun.py @@ -58,7 +58,6 @@ from airflow.models.deadline_alert import DeadlineAlert as DeadlineAlertModel from airflow.models.serialized_dag import SerializedDagModel from airflow.models.taskinstance import TaskInstance, TaskInstanceNote, clear_task_instances -from airflow.models.taskmap import TaskMap from airflow.models.taskreschedule import TaskReschedule from airflow.models.trigger import Trigger from airflow.models.variable import Variable @@ -91,7 +90,7 @@ from tests_common.test_utils.config import conf_vars from tests_common.test_utils.dag import sync_dag_to_db from tests_common.test_utils.db import clear_db_dags, clear_db_runs -from tests_common.test_utils.mapping import expand_mapped_task +from tests_common.test_utils.mapping import expand_mapped_task, push_mapped_length from tests_common.test_utils.mock_operators import MockOperator from tests_common.test_utils.taskinstance import create_task_instance, run_task_instance from unit.models import DEFAULT_DATE as _DEFAULT_DATE @@ -1757,7 +1756,7 @@ def _registered_mutation_hook(hook): """Register hook as the real task_instance_mutation_hook on the policy plugin manager. Patching at the plugin-manager level (rather than airflow.settings) ensures both call sites - see it: TaskMap.expand_mapped_task resolves the wrapper lazily, while refresh_from_task + see it: TaskInstance.expand_mapped_task resolves the wrapper lazily, while refresh_from_task holds a module-level reference bound at import time. """ with mock.patch.object( @@ -1770,7 +1769,7 @@ def _registered_mutation_hook(hook): def test_mutation_hook_committing_session_crashes_under_prohibit_commit(dag_maker, session): """A mutation hook that opens a nested committing session crashes mapped expansion under the guard. - This pins the exact scheduler crash path: during mapped-task expansion (TaskMap.expand_mapped_task) + This pins the exact scheduler crash path: during mapped-task expansion (TaskInstance.expand_mapped_task) the hook is invoked while the outer session is wrapped in prohibit_commit. A hook that calls the @provide_session-decorated TaskInstance.get_dagrun() with no session argument reuses the guarded scoped session; the create_session() context manager then commits on exit, tripping the @@ -1839,7 +1838,7 @@ def safe_hook(task_instance, dag_run=None): def test_mutation_hook_deterministic_across_repeated_invocation_during_expansion(dag_maker, session): """A mutation hook may be invoked more than once per TI during expansion; the result must be stable. - TaskMap.expand_mapped_task invokes the hook on the transient TI and again via refresh_from_task + TaskInstance.expand_mapped_task invokes the hook on the transient TI and again via refresh_from_task after session.merge, so a given mapped index is mutated multiple times. This asserts both that the re-invocation really happens (at least one index sees >1 call) and that a deterministic hook -- one that sets queue as a pure function of TI identity -- yields the same persisted value regardless of how @@ -1875,7 +1874,7 @@ def _make_literal_mapped_dagrun(dag_maker, session, *, dag_id, conf=None): """Build a literal-mapped DAG and its running DagRun, returning (dr, dag_version_id). Unlike _make_mapped_dag_for_expansion (which leaves an xcom-mapped task unexpanded so callers - can drive TaskMap.expand_mapped_task by hand), this builds a literal .expand([...]) so that + can drive TaskInstance.expand_mapped_task by hand), this builds a literal .expand([...]) so that create_dagrun materializes the mapped TIs immediately. Callers can then re-invoke the mutation hook on those persisted TIs by calling dr.verify_integrity(...) -- the real scheduler method -- inside their own prohibit_commit guard. @@ -1895,7 +1894,7 @@ def mapped_task(arg): def test_freshly_built_mapped_ti_exposes_dag_run_as_loaded_none(dag_maker, session): """A freshly-built mapped TaskInstance exposes dag_run as loaded-None, not a lazy-load or raise. - TaskMap.expand_mapped_task constructs each expanded TI with TaskInstance(task, run_id=..., ...) + TaskInstance.expand_mapped_task constructs each expanded TI with TaskInstance(task, run_id=..., ...) and invokes the mutation hook on it before it is merged into a session. A conf-routing hook that resolves the DagRun by attribute access (the _resolve_dagrun discipline) relies on ti.dag_run returning None here -- without hitting the DB and without raising DetachedInstanceError -- so it @@ -2185,7 +2184,7 @@ def task_2(arg2): ... assert ti ti.state = TaskInstanceState.SUCCESS # Behave as if TI ran after: Variable.set(key="arg1", value=[1, 2, 3]) - session.add(TaskMap.from_task_instance_xcom(ti, [1, 2, 3])) + push_mapped_length(ti, [1, 2, 3], session=session) session.flush() decision = dr.task_instance_scheduling_decisions(session=session) @@ -2199,7 +2198,7 @@ def task_2(arg2): ... ti = dr.get_task_instance(task_id="task_1", session=session) assert ti # Behave as if we did and re-ran the task: Variable.set(key="arg1", value=[1, 2, 3, 4]) - session.merge(TaskMap.from_task_instance_xcom(ti, [1, 2, 3, 4])) + push_mapped_length(ti, [1, 2, 3, 4], session=session) ti.state = TaskInstanceState.SUCCESS session.flush() @@ -2237,7 +2236,7 @@ def task_2(arg2): ... assert ti ti.state = TaskInstanceState.SUCCESS # Behave as if TI ran after: Variable.set(key="arg1", value=[1, 2, 3]) - session.add(TaskMap.from_task_instance_xcom(ti, [1, 2, 3])) + push_mapped_length(ti, [1, 2, 3], session=session) session.flush() dr.task_instance_scheduling_decisions(session=session) @@ -2256,7 +2255,7 @@ def task_2(arg2): ... ti = dr.get_task_instance(task_id="task_1", session=session) assert ti # Behave as if we did and re-ran the task: Variable.set(key="arg1", value=[1, 2]) - session.merge(TaskMap.from_task_instance_xcom(ti, [1, 2])) + push_mapped_length(ti, [1, 2], session=session) ti.state = TaskInstanceState.SUCCESS session.flush() dag_version_id = DagVersion.get_latest_version(dag.dag_id, session=session).id @@ -2326,7 +2325,7 @@ def task_2(arg2): ... # "Run" task_1 ti.state = TaskInstanceState.SUCCESS # Behave as if TI ran after: Variable.set(key="arg1", value=[1, 2, 3]) - session.add(TaskMap.from_task_instance_xcom(ti, [1, 2, 3])) + push_mapped_length(ti, [1, 2, 3], session=session) session.flush() decision = dr.task_instance_scheduling_decisions(session=session) @@ -2346,7 +2345,7 @@ def task_2(arg2): ... # We don't execute task anymore, but this is what we are # simulating happened: # Variable.set(key="arg1", value=[]) - session.merge(TaskMap.from_task_instance_xcom(ti, [])) + push_mapped_length(ti, [], session=session) session.flush() # Run the first task again to get the new lengths @@ -2477,9 +2476,7 @@ def test_ti_scheduling_mapped_zero_length(dag_maker, session): dr: DagRun = dag_maker.create_dagrun() ti1, ti2 = sorted(dr.task_instances, key=lambda ti: ti.task_id) ti1.state = TaskInstanceState.SUCCESS - session.add( - TaskMap(dag_id=dr.dag_id, task_id=ti1.task_id, run_id=dr.run_id, map_index=-1, length=0, keys=None) - ) + push_mapped_length(ti1, [], session=session) session.flush() decision = dr.task_instance_scheduling_decisions(session=session) @@ -2566,7 +2563,7 @@ def _task_ids(tis): # After make_list is run, double is expanded. ti = decision.schedulable_tis[0] ti.state = TaskInstanceState.SUCCESS - session.add(TaskMap.from_task_instance_xcom(ti, [1, 2])) + push_mapped_length(ti, [1, 2], session=session) session.flush() decision = dr.task_instance_scheduling_decisions(session=session) @@ -3245,11 +3242,11 @@ def tg(x, y): ("tg.task_2", -1, None), } - # Simulate task_1 execution to produce TaskMap. + # Simulate task_1 execution to produce the mapped length. (ti_1,) = decision.schedulable_tis assert ti_1.task_id == "task_1" ti_1.state = TaskInstanceState.SUCCESS - session.add(TaskMap.from_task_instance_xcom(ti_1, ["a", "b"])) + push_mapped_length(ti_1, ["a", "b"], session=session) session.flush() # Now task_2 in mapped tagk group is expanded. @@ -3611,7 +3608,7 @@ def _task_ids(tis): assert _task_ids(decision.schedulable_tis) == [("push", -1)] ti = decision.schedulable_tis[0] ti.state = TaskInstanceState.SUCCESS - session.add(TaskMap.from_task_instance_xcom(ti, push.function())) + push_mapped_length(ti, push.function(), session=session) session.flush() decision = dr.task_instance_scheduling_decisions(session=session) diff --git a/airflow-core/tests/unit/models/test_mappedoperator.py b/airflow-core/tests/unit/models/test_mappedoperator.py index 3539428be43c3..e34f23f69310e 100644 --- a/airflow-core/tests/unit/models/test_mappedoperator.py +++ b/airflow-core/tests/unit/models/test_mappedoperator.py @@ -30,7 +30,6 @@ from airflow.exceptions import AirflowSkipException from airflow.models.dag_version import DagVersion from airflow.models.taskinstance import TaskInstance -from airflow.models.taskmap import TaskMap from airflow.providers.standard.operators.python import PythonOperator from airflow.sdk import DAG, BaseOperator, TaskGroup, setup, task, task_group, teardown from airflow.serialization.definitions.baseoperator import SerializedBaseOperator @@ -38,7 +37,11 @@ from airflow.utils.state import TaskInstanceState from tests_common.test_utils.dag import sync_dag_to_db -from tests_common.test_utils.mapping import expand_mapped_task +from tests_common.test_utils.mapping import ( + expand_mapped_task, + expand_mapped_task_instances, + push_mapped_length, +) from tests_common.test_utils.mock_operators import MockOperator from tests_common.test_utils.taskinstance import run_task_instance from unit.models import DEFAULT_DATE @@ -112,16 +115,7 @@ def test_expand_mapped_task_instance(dag_maker, session, num_existing_tis, expec dr = dag_maker.create_dagrun() - session.add( - TaskMap( - dag_id=dr.dag_id, - task_id=task1.task_id, - run_id=dr.run_id, - map_index=-1, - length=len(literal), - keys=None, - ) - ) + push_mapped_length(dr.get_task_instance(task1.task_id, session=session), literal, session=session) if num_existing_tis: # Remove the map_index=-1 TI when we're creating other TIs @@ -147,7 +141,7 @@ def test_expand_mapped_task_instance(dag_maker, session, num_existing_tis, expec session.add(ti) session.flush() - TaskMap.expand_mapped_task(mapped_deser, dr.run_id, session=session) + expand_mapped_task_instances(mapped_deser, dr.run_id, session=session) indices = session.execute( select(TaskInstance.map_index, TaskInstance.state) @@ -176,16 +170,7 @@ def test_expand_mapped_task_failed_state_in_db(dag_maker, session): dr = dag_maker.create_dagrun() mapped_deser = dag.task_dict[mapped.task_id] - session.add( - TaskMap( - dag_id=dr.dag_id, - task_id=task1.task_id, - run_id=dr.run_id, - map_index=-1, - length=len(literal), - keys=None, - ) - ) + push_mapped_length(dr.get_task_instance(task1.task_id, session=session), literal, session=session) dag_version = DagVersion.get_latest_version(dr.dag_id) for index in range(2): # Give the existing TIs a state to make sure we don't change them @@ -211,7 +196,7 @@ def test_expand_mapped_task_failed_state_in_db(dag_maker, session): # Make sure we have the faulty state in the database assert indices == [(-1, None), (0, "success"), (1, "success")] - TaskMap.expand_mapped_task(mapped_deser, dr.run_id, session=session) + expand_mapped_task_instances(mapped_deser, dr.run_id, session=session) indices = session.execute( select(TaskInstance.map_index, TaskInstance.state) @@ -278,16 +263,7 @@ def test_expand_kwargs_mapped_task_instance(dag_maker, session, num_existing_tis dr = dag_maker.create_dagrun() - session.add( - TaskMap( - dag_id=dr.dag_id, - task_id=task1.task_id, - run_id=dr.run_id, - map_index=-1, - length=len(literal), - keys=None, - ) - ) + push_mapped_length(dr.get_task_instance(task1.task_id, session=session), literal, session=session) if num_existing_tis: # Remove the map_index=-1 TI when we're creating other TIs @@ -312,7 +288,7 @@ def test_expand_kwargs_mapped_task_instance(dag_maker, session, num_existing_tis session.add(ti) session.flush() - TaskMap.expand_mapped_task(dag.task_dict[mapped.task_id], dr.run_id, session=session) + expand_mapped_task_instances(dag.task_dict[mapped.task_id], dr.run_id, session=session) indices = session.execute( select(TaskInstance.map_index, TaskInstance.state) @@ -349,20 +325,11 @@ def show(number, letter): dr = dag_maker.create_dagrun() for fn in (emit_numbers, emit_letters): - session.add( - TaskMap( - dag_id=dr.dag_id, - task_id=fn.__name__, - run_id=dr.run_id, - map_index=-1, - length=len(fn.function()), - keys=None, - ) - ) + push_mapped_length(dr.get_task_instance(fn.__name__, session=session), fn.function(), session=session) session.flush() show_task = dag.get_task("show") - mapped_tis, max_map_index = TaskMap.expand_mapped_task(show_task, dr.run_id, session=session) + mapped_tis, max_map_index = expand_mapped_task_instances(show_task, dr.run_id, session=session) assert max_map_index + 1 == len(mapped_tis) == 6 diff --git a/airflow-core/tests/unit/models/test_renderedtifields.py b/airflow-core/tests/unit/models/test_renderedtifields.py index f695c46aac274..c6d6809171557 100644 --- a/airflow-core/tests/unit/models/test_renderedtifields.py +++ b/airflow-core/tests/unit/models/test_renderedtifields.py @@ -35,7 +35,6 @@ from airflow.configuration import conf from airflow.models import DagRun from airflow.models.renderedtifields import RenderedTaskInstanceFields as RTIF -from airflow.models.taskmap import TaskMap from airflow.providers.standard.operators.bash import BashOperator from airflow.providers.standard.operators.python import PythonOperator from airflow.sdk import task as task_decorator @@ -44,6 +43,7 @@ from tests_common.test_utils.asserts import assert_queries_count from tests_common.test_utils.db import clear_db_dags, clear_db_runs, clear_rendered_ti_fields +from tests_common.test_utils.mapping import expand_mapped_task_instances if TYPE_CHECKING: from airflow.models.taskinstance import TaskInstance @@ -263,7 +263,7 @@ def test_delete_old_records_mapped( run_id=f"run_{num}", logical_date=dag.start_date + timedelta(days=num) ) - TaskMap.expand_mapped_task(dag.task_dict[mapped.task_id], dr.run_id, session=dag_maker.session) + expand_mapped_task_instances(dag.task_dict[mapped.task_id], dr.run_id, session=dag_maker.session) session.refresh(dr) for ti in dr.task_instances: ti.task = dag_maker.serialized_dag.get_task(ti.task_id) diff --git a/airflow-core/tests/unit/models/test_taskinstance.py b/airflow-core/tests/unit/models/test_taskinstance.py index 7401c7781f3cc..4114378a25fb1 100644 --- a/airflow-core/tests/unit/models/test_taskinstance.py +++ b/airflow-core/tests/unit/models/test_taskinstance.py @@ -71,7 +71,6 @@ find_relevant_relatives, ) from airflow.models.taskinstancehistory import TaskInstanceHistory -from airflow.models.taskmap import TaskMap from airflow.models.taskreschedule import TaskReschedule from airflow.models.xcom import XComModel from airflow.providers.standard.operators.bash import BashOperator @@ -116,6 +115,7 @@ from tests_common.test_utils.asserts import assert_queries_count, capture_orm_selects from tests_common.test_utils.config import conf_vars from tests_common.test_utils.db import clear_db_runs +from tests_common.test_utils.mapping import expand_mapped_task_instances from tests_common.test_utils.mock_operators import MockOperator from tests_common.test_utils.taskinstance import ( create_task_instance as _create_task_instance, @@ -3129,13 +3129,19 @@ def test_noload_relationships_raise_without_joinedload(self, dag_maker, session, getattr(loaded_ti, attr) -class TestTaskInstanceRecordTaskMapXComPush: - """Test TI.xcom_push() correctly records return values for task-mapping.""" +def _mapped_length_count(session) -> int: + return session.scalar( + select(func.count()).select_from(XComModel).where(XComModel.mapped_length.is_not(None)) + ) + + +class TestTaskInstanceRecordMappedLengthXComPush: + """Test TI.xcom_push() correctly records return value lengths for task-mapping.""" def setup_class(self): """Ensure we start fresh.""" with create_session() as session: - session.execute(delete(TaskMap)) + session.execute(delete(XComModel)) @pytest.mark.parametrize("xcom_value", [[1, 2, 3], {"a": 1, "b": 2}, "abc"]) def test_not_recorded_if_leaf(self, dag_maker, xcom_value): @@ -3151,7 +3157,7 @@ def push_something(): ti = next(ti for ti in dag_maker.create_dagrun().task_instances if ti.task_id == "push_something") run_task_instance(ti, dag.get_task(ti.task_id)) - assert dag_maker.session.scalar(select(func.count()).select_from(TaskMap)) == 0 + assert _mapped_length_count(dag_maker.session) == 0 @pytest.mark.parametrize("xcom_value", [[1, 2, 3], {"a": 1, "b": 2}, "abc"]) def test_not_recorded_if_not_used(self, dag_maker, xcom_value): @@ -3171,7 +3177,7 @@ def completely_different(): ti = next(ti for ti in dag_maker.create_dagrun().task_instances if ti.task_id == "push_something") run_task_instance(ti, dag.get_task(ti.task_id)) - assert dag_maker.session.scalar(select(func.count()).select_from(TaskMap)) == 0 + assert _mapped_length_count(dag_maker.session) == 0 @pytest.mark.parametrize("xcom_1", [[1, 2, 3], {"a": 1, "b": 2}, "abc"]) @pytest.mark.parametrize("xcom_4", [[1, 2, 3], {"a": 1, "b": 2}]) @@ -3210,16 +3216,16 @@ def tg(arg): dr = dag_maker.create_dagrun() dag_maker.run_ti("push_1", dr) - assert dag_maker.session.scalar(select(func.count()).select_from(TaskMap)) == 0 + assert _mapped_length_count(dag_maker.session) == 0 dag_maker.run_ti("push_2", dr) - assert dag_maker.session.scalar(select(func.count()).select_from(TaskMap)) == 1 + assert _mapped_length_count(dag_maker.session) == 1 dag_maker.run_ti("push_3", dr) - assert dag_maker.session.scalar(select(func.count()).select_from(TaskMap)) == 1 + assert _mapped_length_count(dag_maker.session) == 1 dag_maker.run_ti("push_4", dr) - assert dag_maker.session.scalar(select(func.count()).select_from(TaskMap)) == 2 + assert _mapped_length_count(dag_maker.session) == 2 class TestMappedTaskInstanceReceiveValue: @@ -3282,7 +3288,7 @@ def show(value): dag_maker.run_ti(emit_ti.task_id, dag_run=dag_run, session=session) show_task = dag_maker.serialized_dag.get_task("show") - mapped_tis, max_map_index = TaskMap.expand_mapped_task(show_task, dag_run.run_id, session=session) + mapped_tis, max_map_index = expand_mapped_task_instances(show_task, dag_run.run_id, session=session) assert max_map_index + 1 == len(mapped_tis) == len(upstream_return) for ti in sorted(mapped_tis, key=operator.attrgetter("map_index")): @@ -3293,7 +3299,7 @@ def show(value): def test_map_xcom_wide_batched_expand(self, dag_maker, session): """Wide XCom-driven expand goes through the batched add_all()/flush() path. - Exercises ``TaskMap.expand_mapped_task`` over a 20-element upstream XCom and + Exercises ``TaskInstance.expand_mapped_task`` over a 20-element upstream XCom and asserts the batched expansion creates exactly N mapped TIs with contiguous ``map_index`` 0..N-1, the expected ``None`` (schedulable) state, and that the returned instances are usable: they keep their ``.task`` (no merge() that drops @@ -3330,7 +3336,9 @@ def show(value): # a slower one. Measured at 7 for this fixture; margin allows for minor backend # differences while staying far below what a per-index merge() would cost. with assert_queries_count(7, margin=2): - mapped_tis, max_map_index = TaskMap.expand_mapped_task(show_task, dag_run.run_id, session=session) + mapped_tis, max_map_index = expand_mapped_task_instances( + show_task, dag_run.run_id, session=session + ) # Correct count + contiguous indexes 0..N-1. assert len(mapped_tis) == width @@ -3406,29 +3414,20 @@ def show(value): dag_maker.run_ti(emit_ti.task_id, dag_run=dag_run, session=session) show_task = dag_maker.serialized_dag.get_task("show") - mapped_tis, max_map_index = TaskMap.expand_mapped_task(show_task, dag_run.run_id, session=session) + mapped_tis, max_map_index = expand_mapped_task_instances(show_task, dag_run.run_id, session=session) assert len(mapped_tis) == 3 assert max_map_index == 2 - # Grow the upstream's pushed length 3 -> 5 by rewriting the return-value XCom and - # the TaskMap row that records the mapped length. + # Grow the upstream's pushed length 3 -> 5 by rewriting the return-value XCom. XComModel.set( key="return_value", value=[1, 2, 3, 4, 5], dag_id=dag_run.dag_id, task_id="emit", run_id=dag_run.run_id, + mapped_length=5, session=session, ) - task_map = session.scalars( - select(TaskMap).where( - TaskMap.dag_id == dag_run.dag_id, - TaskMap.task_id == "emit", - TaskMap.run_id == dag_run.run_id, - ) - ).one() - task_map.length = 5 - task_map.keys = None session.flush() # Pins the query count so a regression back to per-index session.merge() -- which @@ -3470,7 +3469,7 @@ def show(a, b): show_task = dag.get_task("show") assert show_task.get_parse_time_mapped_ti_count() == 6 - mapped_tis, max_map_index = TaskMap.expand_mapped_task(show_task, dag_run.run_id, session=session) + mapped_tis, max_map_index = expand_mapped_task_instances(show_task, dag_run.run_id, session=session) assert len(mapped_tis) == 0 # Expanded at parse! assert max_map_index == 5 @@ -3517,7 +3516,7 @@ def cmds(): dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session) bash_task = dag.get_task("dynamic.bash") - mapped_bash_tis, max_map_index = TaskMap.expand_mapped_task( + mapped_bash_tis, max_map_index = expand_mapped_task_instances( bash_task, dag_run.run_id, session=session ) assert max_map_index == 3 # 2 * 2 mapped tasks. diff --git a/airflow-core/tests/unit/models/test_taskmap.py b/airflow-core/tests/unit/models/test_taskmap.py deleted file mode 100644 index d6a902432268a..0000000000000 --- a/airflow-core/tests/unit/models/test_taskmap.py +++ /dev/null @@ -1,79 +0,0 @@ -# -# 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 unittest import mock - -import pytest - -from airflow.models.taskmap import TaskMap, TaskMapVariant -from airflow.providers.standard.operators.empty import EmptyOperator - -from tests_common.test_utils.taskinstance import create_task_instance - -pytestmark = pytest.mark.db_test - - -def test_task_map_from_task_instance_xcom(): - task = EmptyOperator(task_id="test_task") - ti = create_task_instance(task=task, run_id="test_run", map_index=0, dag_version_id=mock.MagicMock()) - ti.dag_id = "test_dag" - value = {"key1": "value1", "key2": "value2"} - - # Test case where run_id is not None - task_map = TaskMap.from_task_instance_xcom(ti, value) - assert task_map.dag_id == ti.dag_id - assert task_map.task_id == ti.task_id - assert task_map.run_id == ti.run_id - assert task_map.map_index == ti.map_index - assert task_map.length == len(value) - assert task_map.keys == list(value) - - # Test case where run_id is None - ti.run_id = None - with pytest.raises(ValueError, match="cannot record task map for unrun task instance"): - TaskMap.from_task_instance_xcom(ti, value) - - -def test_task_map_with_invalid_task_instance(): - task = EmptyOperator(task_id="test_task") - ti = create_task_instance(task=task, run_id=None, map_index=0, dag_version_id=mock.MagicMock()) - ti.dag_id = "test_dag" - - # Define some arbitrary XCom-like value data - value = {"example_key": "example_value"} - - with pytest.raises(ValueError, match="cannot record task map for unrun task instance"): - TaskMap.from_task_instance_xcom(ti, value) - - -def test_task_map_variant(): - # Test case where keys is None - task_map = TaskMap( - dag_id="test_dag", - task_id="test_task", - run_id="test_run", - map_index=0, - length=3, - keys=None, - ) - assert task_map.variant == TaskMapVariant.LIST - - # Test case where keys is not None - task_map.keys = ["key1", "key2"] - assert task_map.variant == TaskMapVariant.DICT diff --git a/airflow-core/tests/unit/models/test_xcom_arg.py b/airflow-core/tests/unit/models/test_xcom_arg.py index ee29354444932..21f13c5071080 100644 --- a/airflow-core/tests/unit/models/test_xcom_arg.py +++ b/airflow-core/tests/unit/models/test_xcom_arg.py @@ -18,9 +18,12 @@ import pytest +from airflow.models.expandinput import NotFullyPopulated +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 tests_common.test_utils.db import clear_db_dags, clear_db_runs @@ -221,3 +224,45 @@ def pull(value): dag_maker.run_ti(task_id=ti.task_id, map_index=ti.map_index, dag_run=dr, session=session) assert results == expected_results + + +def test_mapped_length_dies_with_the_pushed_value(dag_maker, session): + """Purging the producer's XCom must take its mapped length with it. + + The retry-time purge is issued worker-side so custom XCom backends are purged too; it + deletes XCom rows only. With the length on the row, a downstream expanding while the + producer retries correctly sees "not ready" instead of a length for a value that is gone. + """ + with dag_maker(session=session, serialized=True) as dag: + + @dag.task + def emit(): + return [1, 2, 3] + + @dag.task + def consume(value): ... + + consume.expand(value=emit()) + + dr = dag_maker.create_dagrun() + dag_maker.run_ti(task_id="emit", dag_run=dr, session=session) + session.commit() + + consume_task = dag_maker.serialized_dag.get_task("consume") + assert get_mapped_ti_count(consume_task, dr.run_id, session=session) == 3 + + XComModel.clear(dag_id=dr.dag_id, task_id="emit", run_id=dr.run_id, map_index=-1, session=session) + + with pytest.raises(NotFullyPopulated): + get_mapped_ti_count(consume_task, dr.run_id, session=session) + + XComModel.set( + key=XCOM_RETURN_KEY, + value=[1, 2], + dag_id=dr.dag_id, + task_id="emit", + run_id=dr.run_id, + mapped_length=2, + session=session, + ) + assert get_mapped_ti_count(consume_task, dr.run_id, session=session) == 2 diff --git a/airflow-core/tests/unit/ti_deps/deps/test_mapped_task_upstream_dep.py b/airflow-core/tests/unit/ti_deps/deps/test_mapped_task_upstream_dep.py index 385b6cfeae852..d4d7ad7ee2aa8 100644 --- a/airflow-core/tests/unit/ti_deps/deps/test_mapped_task_upstream_dep.py +++ b/airflow-core/tests/unit/ti_deps/deps/test_mapped_task_upstream_dep.py @@ -21,7 +21,6 @@ import pytest -from airflow.models.taskmap import TaskMap from airflow.models.xcom import XCOM_RETURN_KEY from airflow.providers.standard.operators.empty import EmptyOperator from airflow.sdk.exceptions import AirflowFailException, AirflowSkipException @@ -30,6 +29,8 @@ from airflow.ti_deps.deps.mapped_task_upstream_dep import MappedTaskUpstreamDep from airflow.utils.state import TaskInstanceState +from tests_common.test_utils.mapping import push_mapped_length + pytestmark = [pytest.mark.db_test, pytest.mark.need_serialized_dag] if TYPE_CHECKING: @@ -224,8 +225,7 @@ def tg(x, y): # Simulate running the first schedulable task: t1 returns [0] schedulable_tis["t1"].state = SUCCESS - schedulable_tis["t1"].xcom_push(XCOM_RETURN_KEY, [0], session=session) - session.add(TaskMap.from_task_instance_xcom(schedulable_tis["t1"], [0])) + push_mapped_length(schedulable_tis["t1"], [0], session=session) session.flush() schedulable_tis, finished_tis_states = _one_scheduling_decision_iteration(dr, session) assert sorted(schedulable_tis) == ["t2_a", "t3", "t4"] @@ -242,8 +242,7 @@ def tg(x, y): schedulable_tis, _ = _one_scheduling_decision_iteration(dr, session) if not failure_mode: schedulable_tis["t2_b"].state = SUCCESS - schedulable_tis["t2_b"].xcom_push(XCOM_RETURN_KEY, [1, 2], session=session) - session.add(TaskMap.from_task_instance_xcom(schedulable_tis["t2_b"], [1, 2])) + push_mapped_length(schedulable_tis["t2_b"], [1, 2], session=session) else: schedulable_tis["t2_b"].state = FAILED session.flush() @@ -251,8 +250,7 @@ def tg(x, y): schedulable_tis["t3"].state = SKIPPED else: schedulable_tis["t3"].state = SUCCESS - schedulable_tis["t3"].xcom_push(XCOM_RETURN_KEY, [3, 4], session=session) - session.add(TaskMap.from_task_instance_xcom(schedulable_tis["t3"], [3, 4])) + push_mapped_length(schedulable_tis["t3"], [3, 4], session=session) schedulable_tis["t4"].state = SUCCESS session.flush() _one_scheduling_decision_iteration(dr, session) diff --git a/airflow-core/tests/unit/utils/test_db_cleanup.py b/airflow-core/tests/unit/utils/test_db_cleanup.py index b55ef1ace1b02..c7104d954d43a 100644 --- a/airflow-core/tests/unit/utils/test_db_cleanup.py +++ b/airflow-core/tests/unit/utils/test_db_cleanup.py @@ -1062,7 +1062,6 @@ def test_no_models_missing(self): "asset_active", # not good way to know if "stale" "asset", # not good way to know if "stale" "asset_alias", # not good way to know if "stale" - "task_map", # keys to TI, so no need "serialized_dag", # handled through FK to Dag "log_template", # not a significant source of data; age not indicative of staleness "dag_tag", # not a significant source of data; age not indicative of staleness, diff --git a/devel-common/src/tests_common/pytest_plugin.py b/devel-common/src/tests_common/pytest_plugin.py index 529b4fabc95a4..fbf6aae8b8f47 100644 --- a/devel-common/src/tests_common/pytest_plugin.py +++ b/devel-common/src/tests_common/pytest_plugin.py @@ -1571,7 +1571,6 @@ def __call__( def cleanup(self): from airflow.models import DagModel, DagRun, TaskInstance from airflow.models.serialized_dag import SerializedDagModel - from airflow.models.taskmap import TaskMap from airflow.utils.retries import run_with_db_retries from tests_common.test_utils.compat import AssetEvent @@ -1607,7 +1606,6 @@ def cleanup(self): self.session.execute(delete(TaskInstance).where(TaskInstance.dag_id.in_(dag_ids))) self.session.execute(delete(XCom).where(XCom.dag_id.in_(dag_ids))) self.session.execute(delete(DagModel).where(DagModel.dag_id.in_(dag_ids))) - self.session.execute(delete(TaskMap).where(TaskMap.dag_id.in_(dag_ids))) self.session.execute(delete(AssetEvent).where(AssetEvent.source_dag_id.in_(dag_ids))) if AIRFLOW_V_3_0_PLUS: for bundle_name in self.created_bundle_names: diff --git a/devel-common/src/tests_common/test_utils/mapping.py b/devel-common/src/tests_common/test_utils/mapping.py index b63fadfa7930a..2669f0a3bd03b 100644 --- a/devel-common/src/tests_common/test_utils/mapping.py +++ b/devel-common/src/tests_common/test_utils/mapping.py @@ -18,14 +18,84 @@ from typing import TYPE_CHECKING -from airflow.models.taskmap import TaskMap +from sqlalchemy import select + +from airflow.models.taskinstance import TaskInstance +from airflow.models.xcom import XCOM_RETURN_KEY, XComModel + +from tests_common.test_utils.version_compat import AIRFLOW_V_3_4_PLUS if TYPE_CHECKING: + from collections.abc import Collection, Sequence + from sqlalchemy.orm import Session from airflow.serialization.definitions.mappedoperator import Operator +def push_mapped_length(ti: TaskInstance, value: Collection, *, session: Session) -> None: + """Record ``value`` as ``ti``'s return value, usable as an expansion input.""" + # A dict must keep its shape: expansion hands a mapped task the ``(key, value)`` tuple. + stored = value if isinstance(value, (list, dict)) else list(value) + if AIRFLOW_V_3_4_PLUS: + XComModel.set( + key=XCOM_RETURN_KEY, + value=stored, + dag_id=ti.dag_id, + task_id=ti.task_id, + run_id=ti.run_id, + map_index=ti.map_index, + mapped_length=len(value), + session=session, + ) + return + + # Local import: the module is gone from 3.4 on. + from airflow.models.taskmap import TaskMap + + XComModel.set( + key=XCOM_RETURN_KEY, + value=stored, + dag_id=ti.dag_id, + task_id=ti.task_id, + run_id=ti.run_id, + map_index=ti.map_index, + session=session, + ) + session.add( + TaskMap( + dag_id=ti.dag_id, + task_id=ti.task_id, + run_id=ti.run_id, + map_index=ti.map_index, + length=len(value), + keys=None, + ) + ) + session.flush() + + +def expand_mapped_task_instances( + mapped: Operator, + run_id: str, + *, + session: Session, +) -> tuple[Sequence[TaskInstance], int]: + # -1 sorts first, so this picks the unmapped TI where a test still has one. + ti = session.scalars( + select(TaskInstance) + .where( + TaskInstance.dag_id == mapped.dag_id, + TaskInstance.task_id == mapped.task_id, + TaskInstance.run_id == run_id, + ) + .order_by(TaskInstance.map_index) + .limit(1) + ).one() + ti.task = mapped + return ti.expand_mapped_task(session=session) + + def expand_mapped_task( mapped: Operator, run_id: str, @@ -33,16 +103,13 @@ def expand_mapped_task( length: int, session: Session, ): - session.add( - TaskMap( - dag_id=mapped.dag_id, - task_id=upstream_task_id, - run_id=run_id, - map_index=-1, - length=length, - keys=None, + upstream_ti = session.scalars( + select(TaskInstance).where( + TaskInstance.dag_id == mapped.dag_id, + TaskInstance.task_id == upstream_task_id, + TaskInstance.run_id == run_id, + TaskInstance.map_index == -1, ) - ) - session.flush() - - TaskMap.expand_mapped_task(mapped, run_id, session=session) + ).one() + push_mapped_length(upstream_ti, list(range(length)), session=session) + expand_mapped_task_instances(mapped, run_id, session=session) diff --git a/providers/standard/tests/unit/standard/decorators/test_python.py b/providers/standard/tests/unit/standard/decorators/test_python.py index fb972a652d443..252ecc5c1b059 100644 --- a/providers/standard/tests/unit/standard/decorators/test_python.py +++ b/providers/standard/tests/unit/standard/decorators/test_python.py @@ -23,7 +23,6 @@ import pytest -from airflow.models.taskmap import TaskMap from airflow.providers.common.compat.sdk import AirflowException, XComNotFound from tests_common.test_utils.taskinstance import get_template_context, render_template_fields @@ -833,6 +832,8 @@ def task2(arg1, arg2): ... def test_mapped_render_template_fields(dag_maker, session): from airflow.sdk.definitions.mappedoperator import MappedOperator + from tests_common.test_utils.mapping import push_mapped_length + @task_decorator def fn(arg1, arg2): ... @@ -843,18 +844,7 @@ def fn(arg1, arg2): ... dr = dag_maker.create_dagrun() ti: TaskInstance = dr.get_task_instance(task1.task_id, session=session) - ti.xcom_push(key=XCOM_RETURN_KEY, value=["{{ ds }}"], session=session) - - session.add( - TaskMap( - dag_id=dr.dag_id, - task_id=task1.task_id, - run_id=dr.run_id, - map_index=-1, - length=1, - keys=None, - ) - ) + push_mapped_length(ti, ["{{ ds }}"], session=session) session.flush() mapped_ti: TaskInstance = dr.get_task_instance(mapped.operator.task_id, session=session) @@ -871,6 +861,7 @@ def fn(arg1, arg2): ... @pytest.mark.skipif(AIRFLOW_V_3_0_PLUS, reason="Different test for AF 2") def test_mapped_render_template_fields_af2(dag_maker, session): + from airflow.models.taskmap import TaskMap from airflow.utils.task_instance_session import set_current_task_instance_session @task_decorator diff --git a/scripts/cov/core_coverage.py b/scripts/cov/core_coverage.py index b7096335a4300..cf7cf49f8e402 100644 --- a/scripts/cov/core_coverage.py +++ b/scripts/cov/core_coverage.py @@ -63,7 +63,6 @@ "airflow-core/src/airflow/models/taskinstance.py", "airflow-core/src/airflow/models/taskinstancehistory.py", "airflow-core/src/airflow/models/taskinstancekey.py", - "airflow-core/src/airflow/models/taskmap.py", "airflow-core/src/airflow/models/taskmixin.py", "airflow-core/src/airflow/models/trigger.py", "airflow-core/src/airflow/models/variable.py",