Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
42 commits
Select commit Hold shift + click to select a range
01d5518
AIP-104: Task Iteration
dabla Oct 8, 2026
4ec2b3c
Describe the behaviour in two iteration test docstrings
dabla Oct 8, 2026
1e420d2
Hand an item its own operator as the context's task
dabla Oct 8, 2026
264ab01
Say that item checkpoints keep results in the task state store
dabla Oct 8, 2026
7c4f49c
Document that items emitting to one asset share one event
dabla Oct 8, 2026
c51ff17
Refuse an unmappable upstream value as an iteration input
dabla Oct 8, 2026
9ed0ec6
Keep the retry policy's decision for the exception handed to the runner
dabla Oct 8, 2026
0e027c3
Run a sync item's enter and exit in its worker thread, and its kill o…
dabla Oct 8, 2026
6791145
Stop the iteration once the iterated task is killed
dabla Oct 8, 2026
f419f42
Keep the state of one iteration run in IterationState
dabla Oct 8, 2026
9f96f41
Keep default_args callbacks and execute hooks with the items
dabla Oct 8, 2026
2756e26
Kill in a thread of its own whenever on_kill() is called
dabla Oct 8, 2026
96b14f2
Kill the operator the parent's timeout strikes off the loop thread, a…
dabla Oct 8, 2026
bfe5db5
Send the parent's execution timeout to the supervisor once, not per s…
dabla Oct 8, 2026
ec4febb
Say that listeners fire once per iterated task instance, not per item
dabla Oct 9, 2026
9eadec1
Read an iteration's own XComs back under their indexed key
dabla Oct 9, 2026
24551fa
Classify an item's outcome in one method
dabla Oct 9, 2026
c1b0068
Keep what the indexed tasks ended with in IndexedTaskOutcomes
dabla Oct 9, 2026
c908268
Give the iteration's helper functions to the classes that own them
dabla Oct 9, 2026
63f4ee4
Refuse the backend clear on the indexed state store view and keep a k…
dabla Oct 9, 2026
f582e1e
Say iterate in the error a @task raises when iterate() gets no arguments
dabla Oct 9, 2026
c03f879
Let clone_context clone a context without inlet_events or dag_run
dabla Oct 9, 2026
c2c763a
Point async sub-tasks at XCom.aget_one for another iteration's value
dabla Oct 9, 2026
f811f22
Render an indexed task's templates against the same context keys it e…
dabla Oct 9, 2026
6d684eb
Say that an iteration returning None keeps its position and reads bac…
dabla Oct 9, 2026
797aa52
Join the SIGTERM kill thread before the iteration concludes
dabla Oct 9, 2026
4e722a4
Add the newsfragment for Task Iteration
dabla Oct 9, 2026
1e566d4
Report an indexed task's success or skip once its checkpoint is written
dabla Oct 9, 2026
d52ca9c
Keep the default decision for an undecided failure group instead of e…
dabla Oct 9, 2026
7d834e9
Do not start an indexed task instance pulled before the kill reached …
dabla Oct 9, 2026
6df446b
Write the context_for docstring in the imperative mood
dabla Oct 9, 2026
6b51e64
Name only the counts a kill left behind in the terminated task's message
dabla Oct 9, 2026
5f7a4db
Name the workers section the state store backend option lives in
dabla Oct 10, 2026
5019f84
Declare the delegate of a decorated expand input as a field so two in…
dabla Oct 10, 2026
730ef65
Refuse a literal that is no collection in iterate() as expand() does
dabla Oct 10, 2026
eb56027
Report an indexed task's success where a failing report leaves its ch…
dabla Oct 10, 2026
b4ada44
Give the copy prepared for execution an iteration state of its own
dabla Oct 10, 2026
bd177a3
Take no lock in on_kill, which the SIGTERM handler runs on the loop t…
dabla Oct 10, 2026
35a98d4
Treat an AirflowTaskTimeout a sync indexed task raises itself as that…
dabla Oct 10, 2026
e3aeb0c
Drop the per-indexed-task limit that could never fire before the pare…
dabla Oct 10, 2026
616fea8
Release the comms thread lock an asend cancelled while waiting for it…
dabla Oct 10, 2026
8a44bf7
Say when not to iterate an operator that finds its remote work by the…
dabla Oct 10, 2026
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
2 changes: 2 additions & 0 deletions .github/CODEOWNERS
Original file line number Diff line number Diff line change
Expand Up @@ -163,6 +163,8 @@ Dockerfile.ci @potiuk @ashb @gopidesupavan @amoghrajesh @jscheffl @bugraoz93 @ja
# AIP-72 - Task SDK
# Python SDK
/task-sdk/ @ashb @amoghrajesh
/task-sdk/src/airflow/sdk/execution_time/executor.py @ashb @amoghrajesh @dabla
/task-sdk/src/airflow/sdk/definitions/iterableoperator.py @ashb @amoghrajesh @dabla

# AIP-108 - Coordinators
/task-sdk/src/airflow/sdk/coordinators/ @jason810496 @uranusjr
Expand Down
1 change: 1 addition & 0 deletions airflow-core/newsfragments/62922.feature.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Task Iteration (AIP-104): ``.iterate()`` and ``.iterate_kwargs()`` on operators and ``@task`` run the mapped inputs inside one task instance, ``task_concurrency`` of them at a time, instead of one task instance per item as ``.expand()`` does. Each iteration's result is checkpointed in the task state store, so a retry resumes where the previous attempt stopped, and the task's return value is an ``XComIterable``, a lazy sequence downstream tasks index, iterate or ``.expand()`` over. The Task SDK docs page "Mapped tasks vs iterable tasks" says when to use which.
341 changes: 341 additions & 0 deletions airflow-core/tests/unit/models/test_taskinstance.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@
from airflow.exceptions import (
AirflowException,
AirflowSkipException,
NotMapped,
)
from airflow.listeners.listener import get_listener_manager
from airflow.models.asset import (
Expand Down Expand Up @@ -3950,6 +3951,346 @@ def show(a, b):
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
assert outputs == [(2, 5), (2, 10), (4, 5), (4, 10), (8, 5), (8, 10)]

def test_iterate_literal_cross_product(self, dag_maker, session):
"""Test an iterated task with literal cross product args properly."""
outputs = []

with dag_maker(dag_id="product_same_types", session=session, serialized=True) as dag:

@dag.task
def show(a, b):
outputs.append((a, b))

show.iterate(a=[2, 4, 8], b=[5, 10])

dag_run = dag_maker.create_dagrun()

show_task = dag.get_task("show")
with pytest.raises(NotMapped):
show_task.get_parse_time_mapped_ti_count()
with pytest.raises(NotMapped):
expand_mapped_task_instances(show_task, dag_run.run_id, session=session)

tis = session.scalars(
select(TaskInstance)
.where(
TaskInstance.dag_id == dag.dag_id,
TaskInstance.task_id == "show",
TaskInstance.run_id == dag_run.run_id,
)
.order_by(TaskInstance.map_index)
).all()
for ti in tis:
ti.refresh_from_task(show_task)
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
assert outputs == [(2, 5), (2, 10), (4, 5), (4, 10), (8, 5), (8, 10)]

@pytest.mark.parametrize(
("upstream_is_iterated", "trigger_rule"),
[
pytest.param(False, "all_success", id="expand-all_success"),
pytest.param(True, "all_success", id="iterate-all_success"),
pytest.param(False, "none_failed", id="expand-none_failed"),
pytest.param(True, "none_failed", id="iterate-none_failed"),
],
)
def test_skipped_item_downstream_matches_expand(
self, dag_maker, session, upstream_is_iterated, trigger_rule
):
"""
One skipped item of an upstream has the same effect downstream whether it is mapped or iterated.

With ``all_success`` the downstream task is skipped and receives nothing; with ``none_failed``
it runs over the items that produced a value, the skipped one left out.
"""
received = []

with dag_maker(dag_id=f"skipped_item_{trigger_rule}", session=session, serialized=True):

@task
def produce(x):
if x == 2:
raise AirflowSkipException("nothing to do for this item")
return x * 10

@task(trigger_rule=trigger_rule)
def consume(value):
received.append(value)

produced = produce.iterate(x=[1, 2, 3]) if upstream_is_iterated else produce.expand(x=[1, 2, 3])
consume.expand(value=produced)

dag_run = dag_maker.create_dagrun()
for task_id in ("produce", "consume"):
dag_run.refresh_from_db(session=session)
for ti in dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis:
if ti.task_id == task_id:
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
session.flush()
dag_run.refresh_from_db(session=session)
dag_run.task_instance_scheduling_decisions(session=session)
session.flush()

consume_states = {
ti.state
for ti in session.scalars(
select(TaskInstance).where(
TaskInstance.run_id == dag_run.run_id, TaskInstance.task_id == "consume"
)
)
}
if trigger_rule == "all_success":
assert received == []
assert consume_states == {TaskInstanceState.SKIPPED}
else:
assert sorted(received) == [10, 30]
assert consume_states == {TaskInstanceState.SUCCESS}

@pytest.mark.parametrize("upstream_is_iterated", [False, True], ids=["expand", "iterate"])
def test_empty_input_skips_like_expand(self, dag_maker, session, upstream_is_iterated):
"""Over an empty input the task is skipped and an ``all_success`` downstream task with it."""
ran = []

with dag_maker(dag_id="empty_input", session=session, serialized=True):

@task
def produce(x):
ran.append(x)
return x

@task
def consume(values):
ran.append(values)

produced = produce.iterate(x=[]) if upstream_is_iterated else produce.expand(x=[])
consume(produced)

dag_run = dag_maker.create_dagrun()
for task_id in ("produce", "consume"):
dag_run.refresh_from_db(session=session)
for ti in dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis:
if ti.task_id == task_id:
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
session.flush()
dag_run.refresh_from_db(session=session)
dag_run.task_instance_scheduling_decisions(session=session)
session.flush()

states = {
ti.task_id: ti.state
for ti in session.scalars(select(TaskInstance).where(TaskInstance.run_id == dag_run.run_id))
}
assert ran == []
assert states == {"produce": TaskInstanceState.SKIPPED, "consume": TaskInstanceState.SKIPPED}

def test_iterate_in_task_group_reaches_downstream(self, dag_maker, session):
"""An iterated task in a TaskGroup pushes its results for its own task id, where downstream reads them."""
received = []

with dag_maker(dag_id="iterate_in_task_group", session=session, serialized=True) as dag:
with TaskGroup("group"):

@task
def produce(x):
return x * 10

@task
def consume(value):
received.append(value)

consume.expand(value=produce.iterate(x=[1, 2, 3]))

assert sorted(dag.task_dict) == ["group.consume", "group.produce"]
dag_run = dag_maker.create_dagrun()
for task_id in ("group.produce", "group.consume"):
dag_run.refresh_from_db(session=session)
for ti in dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis:
if ti.task_id == task_id:
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
session.flush()

assert sorted(received) == [10, 20, 30]

@pytest.mark.parametrize(
("with_policy", "expected_state"),
[
pytest.param(True, TaskInstanceState.FAILED, id="policy-fails-it"),
pytest.param(False, TaskInstanceState.UP_FOR_RETRY, id="no-policy-retries"),
],
)
def test_iterate_retry_policy_decides_on_the_items_exception(
self, dag_maker, session, with_policy, expected_state
):
"""A retry policy rule on an item's exception decides the iterated task's outcome, as for any task."""
from airflow.sdk import ExceptionRetryPolicy, RetryRule
from airflow.sdk.definitions.retry_policy import RetryAction

policy = ExceptionRetryPolicy(rules=[RetryRule(exception=PermissionError, action=RetryAction.FAIL)])

with dag_maker(dag_id=f"iterate_retry_policy_{with_policy}", session=session, serialized=True):

@task(retries=2, retry_policy=policy if with_policy else None)
def produce(x):
if x == 2:
raise PermissionError("not allowed")
if x == 3:
raise ValueError("flaky")
return x

produce.iterate(x=[1, 2, 3])

dag_run = dag_maker.create_dagrun()
(ti,) = dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis
# run_ti re-raises the task's error once the outcome is recorded; the state is what counts.
with contextlib.suppress(BaseException):
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
ti.refresh_from_db(session=session)

assert ti.state == expected_state

def test_iterate_sensor_timeout_fails_without_retry(self, dag_maker, session):
"""A poke-mode sensor timing out inside .iterate() fails the task, as it does outside."""
from airflow.sdk.exceptions import AirflowSensorTimeout

with dag_maker(dag_id="iterate_sensor_timeout", session=session, serialized=True):

@task(retries=2)
def produce(x):
if x == 2:
raise AirflowSensorTimeout("poked for too long")
raise ValueError("flaky")

produce.iterate(x=[1, 2])

dag_run = dag_maker.create_dagrun()
(ti,) = dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis
# run_ti re-raises the task's error once the outcome is recorded; the state is what counts.
with contextlib.suppress(BaseException):
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
ti.refresh_from_db(session=session)

assert ti.state == TaskInstanceState.FAILED

def test_iterate_retry_keeps_the_extra_xcoms_of_items_that_already_succeeded(self, dag_maker, session):
"""
The runner deletes every XCom of the task before a retry. An item skipped on the retry
because it already succeeded gets its other pushed keys back from its checkpoint.
"""
from airflow.models.xcom import XComModel

attempts = []

with dag_maker(dag_id="iterate_extra_xcoms", session=session, serialized=True):
# One item at a time: several sync items pushing concurrently under dag_maker's in-process
# supervisor race on it, which is a separate problem.
@task(retries=1, retry_delay=datetime.timedelta(0), task_concurrency=1)
def produce(x, ti=None):
ti.xcom_push(key="foo", value=f"foo-of-{x}")
if x == 2 and not attempts:
attempts.append(ti.try_number)
raise ValueError("flaky on the first attempt")
return x

produce.iterate(x=[1, 2])

dag_run = dag_maker.create_dagrun()
for _ in range(2):
dag_run.refresh_from_db(session=session)
for ti in dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis:
with contextlib.suppress(BaseException):
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
session.flush()

ti = dag_run.get_task_instance("produce", session=session)
assert ti.state == TaskInstanceState.SUCCESS
assert ti.try_number == 2
keys = set(
session.scalars(
select(XComModel.key).where(
XComModel.dag_id == "iterate_extra_xcoms",
XComModel.run_id == dag_run.run_id,
XComModel.task_id == "produce",
)
)
)
assert {"foo_0", "foo_1", "return_value_0", "return_value_1"} <= keys

@pytest.mark.xfail(
strict=True,
reason="dag.test()'s in-process supervisor is not safe to call from several threads until apache/airflow#74074",
)
def test_iterate_concurrent_sync_items_share_the_in_process_supervisor(self, dag_maker, session):
"""
Under the in-process supervisor (dag.test(), dag_maker) sync items make SDK calls from worker
threads; they must not race on it: every item's XCom push and read goes through.
"""
from airflow.models.xcom import XComModel

with dag_maker(dag_id="iterate_in_process_comms", session=session, serialized=True):

@task(task_concurrency=4)
def produce(x, ti=None):
for n in range(5):
ti.xcom_push(key=f"k{n}", value=x)
ti.xcom_pull(task_ids="produce", key=f"k{n}_0")
return x

produce.iterate(x=list(range(8)))

dag_run = dag_maker.create_dagrun()
(ti,) = dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis
with contextlib.suppress(BaseException):
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
ti.refresh_from_db(session=session)

assert ti.state == TaskInstanceState.SUCCESS
keys = set(
session.scalars(
select(XComModel.key).where(
XComModel.dag_id == "iterate_in_process_comms",
XComModel.run_id == dag_run.run_id,
XComModel.task_id == "produce",
)
)
)
assert {f"return_value_{i}" for i in range(8)} <= keys
assert {f"k{n}_{i}" for n in range(5) for i in range(8)} <= keys

@pytest.mark.xfail(
strict=True,
reason="dag.test()'s in-process supervisor is not safe to call from several threads until apache/airflow#74074",
)
def test_iterate_concurrent_sync_items_read_variables_while_a_sibling_is_served(self, dag_maker, session):
"""
Under the in-process supervisor, an item's Variable lookup must not miss while a sibling's
request is served: hiding the comms from the whole process sent it to the fallback secrets
backends, which do not have a Variable stored in the metadata database.
"""
from airflow.models.variable import Variable
from airflow.sdk import Variable as SdkVariable

Variable.set(key="iterate_db_variable", value="v", session=session)
session.commit()

with dag_maker(dag_id="iterate_in_process_variables", session=session, serialized=True):

@task(task_concurrency=4)
def read(x, ti=None):
for _ in range(25):
assert SdkVariable.get("iterate_db_variable") == "v"
ti.xcom_push(key="k", value=x)
return x

read.iterate(x=list(range(8)))

dag_run = dag_maker.create_dagrun()
(ti,) = dag_run.task_instance_scheduling_decisions(session=session).schedulable_tis
with contextlib.suppress(BaseException):
dag_maker.run_ti(ti.task_id, map_index=ti.map_index, dag_run=dag_run, session=session)
ti.refresh_from_db(session=session)

assert ti.state == TaskInstanceState.SUCCESS

def test_map_in_group(self, tmp_path: pathlib.Path, dag_maker, session):
out = tmp_path.joinpath("out")
out.touch()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -896,6 +896,8 @@ def validate_deserialized_task(
"partial_kwargs",
"expand_input",
"weight_rule",
# Always True for serialized operators; excluded from serialization.
"_register_with_dag",
}

assert serialized_task.task_type == task.task_type
Expand Down
18 changes: 18 additions & 0 deletions airflow-core/tests/unit/serialization/test_serialized_objects.py
Original file line number Diff line number Diff line change
Expand Up @@ -1821,3 +1821,21 @@ def test_deadline_callback_rejects_field_belonging_to_the_other_subclass(callbac

with pytest.raises(ValueError, match=f"Unexpected deadline callback fields: {foreign_field}"):
BaseSerialization.deserialize(serialized)


def test_operator_iterate_serde_refers_to_the_class_that_runs():
"""The class an iterated task is serialized under can be imported: the IterableOperator that runs it."""
from airflow._shared.module_loading import import_string
from airflow.providers.standard.operators.bash import BashOperator
from airflow.sdk import DAG
from airflow.sdk.definitions.iterableoperator import IterableOperator
from airflow.serialization.serialized_objects import OperatorSerialization

with DAG("iterated"):
real_op = BashOperator.partial(task_id="a").iterate(bash_command=["echo 1"])

serialized = OperatorSerialization.serialize_operator(real_op)

assert import_string(f"{serialized['_task_module']}.{serialized['task_type']}") is IterableOperator
assert serialized["_operator_name"] == "BashOperator"
assert OperatorSerialization.deserialize_operator(serialized).operator_name == "BashOperator"
Loading
Loading