From 0fb08cac1f3430dbb76f6fc724d1956d13c625c3 Mon Sep 17 00:00:00 2001 From: Ash Berlin-Taylor Date: Thu, 10 Sep 2026 19:17:18 +0100 Subject: [PATCH] Remove the task_map table, folding its length into xcom.mapped_length The task_map table was created (by me and TP) years ago to store the number of results (length) that a task upstream of a mapped task created. This is then used by the scheduler to know how many mapped tasks to expand to. The `keys` column has never been read from (and only got written to in 2.x, but not in any 3.x), it was probably to let people iterate over tasks that returned dictionaries, but well, that turned out to not be required. This is a pre-cursor tidy up to make the PRs for task loops (AIP-111) and after that dynamic sub-graphs (AIP-113) simpler. The wire protocol for both public API and Exec API remain unchanged, this storage on the DB table was purely a server-side storage decision. --- airflow-core/docs/migrations-ref.rst | 4 +- .../core_api/routes/public/xcom.py | 2 + .../api_fastapi/core_api/services/ui/grid.py | 5 +- .../api_fastapi/execution_api/routes/xcoms.py | 26 +- ...0_fold_task_map_into_xcom_mapped_length.py | 158 ++++++++++ airflow-core/src/airflow/models/dagrun.py | 3 +- .../src/airflow/models/taskinstance.py | 166 +++++++++- airflow-core/src/airflow/models/taskmap.py | 293 ------------------ airflow-core/src/airflow/models/xcom.py | 9 + .../serialization/definitions/xcom_arg.py | 26 +- airflow-core/src/airflow/utils/db.py | 2 +- .../routes/public/test_task_instances.py | 15 +- .../core_api/routes/public/test_xcom.py | 37 ++- .../versions/head/test_task_instances.py | 6 +- .../execution_api/versions/head/test_xcoms.py | 67 ++-- .../test_0136_fold_task_map_into_xcom.py | 184 +++++++++++ airflow-core/tests/unit/models/test_dagrun.py | 37 +-- .../tests/unit/models/test_mappedoperator.py | 59 +--- .../unit/models/test_renderedtifields.py | 4 +- .../tests/unit/models/test_taskinstance.py | 53 ++-- .../tests/unit/models/test_taskmap.py | 79 ----- .../tests/unit/models/test_xcom_arg.py | 45 +++ .../deps/test_mapped_task_upstream_dep.py | 12 +- .../tests/unit/utils/test_db_cleanup.py | 1 - .../src/tests_common/pytest_plugin.py | 2 - .../src/tests_common/test_utils/mapping.py | 93 +++++- .../unit/standard/decorators/test_python.py | 17 +- scripts/cov/core_coverage.py | 1 - 28 files changed, 830 insertions(+), 576 deletions(-) create mode 100644 airflow-core/src/airflow/migrations/versions/0136_3_4_0_fold_task_map_into_xcom_mapped_length.py delete mode 100644 airflow-core/src/airflow/models/taskmap.py create mode 100644 airflow-core/tests/unit/migrations/test_0136_fold_task_map_into_xcom.py delete mode 100644 airflow-core/tests/unit/models/test_taskmap.py 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",