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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion airflow-core/docs/migrations-ref.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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. |
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]]:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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")
Comment thread
ashb marked this conversation as resolved.
] = None,
):
"""Set an Airflow XCom."""
Expand All @@ -415,24 +414,24 @@ 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={
"reason": "unmappable_return_value_length",
"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
Expand All @@ -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:
Expand Down
Original file line number Diff line number Diff line change
@@ -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")
3 changes: 1 addition & 2 deletions airflow-core/src/airflow/models/dagrun.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
Loading
Loading