diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index cad6d91e6f8ac..bca715e788524 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -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 diff --git a/airflow-core/newsfragments/62922.feature.rst b/airflow-core/newsfragments/62922.feature.rst new file mode 100644 index 0000000000000..934b417784979 --- /dev/null +++ b/airflow-core/newsfragments/62922.feature.rst @@ -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. diff --git a/airflow-core/tests/unit/models/test_taskinstance.py b/airflow-core/tests/unit/models/test_taskinstance.py index 0820c2f788118..50a0900945fb4 100644 --- a/airflow-core/tests/unit/models/test_taskinstance.py +++ b/airflow-core/tests/unit/models/test_taskinstance.py @@ -48,6 +48,7 @@ from airflow.exceptions import ( AirflowException, AirflowSkipException, + NotMapped, ) from airflow.listeners.listener import get_listener_manager from airflow.models.asset import ( @@ -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() diff --git a/airflow-core/tests/unit/serialization/test_dag_serialization.py b/airflow-core/tests/unit/serialization/test_dag_serialization.py index 29b2890b8e3af..1fe3e147b99e1 100644 --- a/airflow-core/tests/unit/serialization/test_dag_serialization.py +++ b/airflow-core/tests/unit/serialization/test_dag_serialization.py @@ -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 diff --git a/airflow-core/tests/unit/serialization/test_serialized_objects.py b/airflow-core/tests/unit/serialization/test_serialized_objects.py index f1d44c211f027..a6d6f5e3ba8f8 100644 --- a/airflow-core/tests/unit/serialization/test_serialized_objects.py +++ b/airflow-core/tests/unit/serialization/test_serialized_objects.py @@ -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" diff --git a/devel-common/src/tests_common/test_utils/mock_context.py b/devel-common/src/tests_common/test_utils/mock_context.py index 4e7aa7884f259..9c34b4f5816bb 100644 --- a/devel-common/src/tests_common/test_utils/mock_context.py +++ b/devel-common/src/tests_common/test_utils/mock_context.py @@ -20,14 +20,28 @@ from typing import TYPE_CHECKING, Any from unittest import mock +from airflow.models import DagRun +from airflow.providers.common.compat.sdk import timezone +from airflow.utils.types import DagRunType + from tests_common.test_utils.compat import Context from tests_common.test_utils.taskinstance import create_task_instance +from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS if TYPE_CHECKING: from sqlalchemy.orm import Session + from airflow.sdk.types import Operator + from airflow.serialization.definitions.mappedoperator import Operator as SerializedOperator + + +def generate_run_id() -> str: + if AIRFLOW_V_3_0_PLUS: + return DagRun.generate_run_id(run_type=DagRunType.MANUAL, run_after=timezone.utcnow()) + return DagRun.generate_run_id(run_type=DagRunType.MANUAL, execution_date=timezone.utcnow()) # type: ignore[call-arg] -def mock_context(task) -> Context: + +def mock_context(task: Operator | SerializedOperator, run_id: str | None = None) -> Context: from airflow.models import TaskInstance from airflow.utils.session import NEW_SESSION @@ -40,6 +54,29 @@ def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.values: dict[str, Any] = {} + async def axcom_pull( + self, + task_ids: str | Iterable[str] | None = None, + dag_id: str | None = None, + key: str = XCOM_RETURN_KEY, + include_prior_dates: bool = False, + session: Session = NEW_SESSION, + *, + map_indexes: int | Iterable[int] | None = None, + default: Any = None, + run_id: str | None = None, + ) -> Any: + return self.xcom_pull( + task_ids=task_ids, + dag_id=dag_id, + key=key, + include_prior_dates=include_prior_dates, + session=session, + map_indexes=map_indexes, + default=default, + run_id=run_id, + ) + def xcom_pull( self, task_ids: str | Iterable[str] | None = None, @@ -57,13 +94,18 @@ def xcom_pull( key += f"_{map_indexes}" return values.get(key, default) + async def axcom_push(self, key: str, value: Any, session: Session = NEW_SESSION, **kwargs) -> None: + self.xcom_push(key=key, value=value, session=session, **kwargs) + def xcom_push(self, key: str, value: Any, session: Session = NEW_SESSION, **kwargs) -> None: key = f"{self.task_id}_{self.dag_id}_{key}" if self.map_index is not None and self.map_index >= 0: key += f"_{self.map_index}" values[key] = value + values["task"] = task values["ti"] = create_task_instance(task, dag_version_id=mock.MagicMock(), ti_type=MockedTaskInstance) + values["task_instance"] = values["ti"] + values["run_id"] = generate_run_id() if run_id is None else run_id - # See https://github.com/python/mypy/issues/8890 - mypy does not support passing typed dict to TypedDict return Context(values) # type: ignore[misc] diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt index 325b68d6c9027..b27ad7c6fa931 100644 --- a/docs/spelling_wordlist.txt +++ b/docs/spelling_wordlist.txt @@ -1299,6 +1299,7 @@ PodManager PodSpec podSpec podspec +Pokémon polars poller polyfill diff --git a/providers/openlineage/tests/unit/openlineage/extractors/test_manager.py b/providers/openlineage/tests/unit/openlineage/extractors/test_manager.py index d153b5b78951c..5d01c42fea3eb 100644 --- a/providers/openlineage/tests/unit/openlineage/extractors/test_manager.py +++ b/providers/openlineage/tests/unit/openlineage/extractors/test_manager.py @@ -45,7 +45,7 @@ from tests_common.test_utils.compat import DateTimeSensor, PythonOperator from tests_common.test_utils.markers import skip_if_force_lowest_dependencies_marker -from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_2_PLUS +from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_2_PLUS, AIRFLOW_V_3_4_PLUS if TYPE_CHECKING: try: @@ -402,6 +402,34 @@ def execute(self, context: Context) -> Any: assert metadata.outputs == outlets +@pytest.mark.skipif(not AIRFLOW_V_3_4_PLUS, reason=".iterate() is available from Airflow 3.4") +def test_extract_metadata_of_an_iterated_operator_falls_back_to_its_inlets_and_outlets(): + """ + An iterated operator is run by an IterableOperator, which has none of the wrapped operator's + attributes: the wrapped operator's extractor must not be picked for it, so the declared + inlets and outlets are emitted rather than dropped on the extractor's AttributeError. + """ + from airflow.providers.standard.operators.bash import BashOperator + from airflow.sdk import DAG + + inlets = [OpenLineageDataset(namespace="namespace1", name="name1")] + outlets = [OpenLineageDataset(namespace="namespace2", name="name2")] + with DAG("iterated_bash"): + task = BashOperator.partial(task_id="bash", inlets=inlets, outlets=outlets).iterate( + bash_command=["echo 1", "echo 2"] + ) + + extractor_manager = ExtractorManager() + assert extractor_manager.get_extractor_class(task) is None + + metadata = extractor_manager.extract_metadata( + dagrun=MagicMock(), task=task, task_instance_state=None, task_instance=MagicMock() + ) + + assert metadata.inputs == inlets + assert metadata.outputs == outlets + + def test_get_extractor_supports_legacy_custom_extractor_signature(): """ Regression: custom extractors may use the historically-public ``__init__(self, operator)`` diff --git a/task-sdk/docs/deferred-vs-async-operators.rst b/task-sdk/docs/deferred-vs-async-operators.rst index 4212e36246109..aa664312b93d5 100644 --- a/task-sdk/docs/deferred-vs-async-operators.rst +++ b/task-sdk/docs/deferred-vs-async-operators.rst @@ -205,11 +205,11 @@ concurrently using ``asyncio.gather`` while limiting concurrency with a semaphor .. note:: - The upcoming *Dynamic Task Iteration* feature will simplify patterns like this. + :ref:`Iterable Tasks (IT) ` simplifies patterns like this. Instead of manually managing concurrency with constructs such as - ``asyncio.gather`` and ``asyncio.Semaphore``, authors will be able to iterate + ``asyncio.gather`` and ``asyncio.Semaphore``, authors can iterate over asynchronous results directly in downstream tasks while still benefiting - from a shared event loop. This will make high-throughput patterns such as + from a shared event loop. This makes high-throughput patterns such as pagination or request multiplexing easier to implement. MS Graph Async Example diff --git a/task-sdk/docs/index.rst b/task-sdk/docs/index.rst index 9e15dae85e550..71d11ac428ce6 100644 --- a/task-sdk/docs/index.rst +++ b/task-sdk/docs/index.rst @@ -175,6 +175,7 @@ For the full public API reference, see the :doc:`api` page. examples dynamic-task-mapping + mapped-tasks-vs-iterable-tasks deferred-vs-async-operators resumable-job-mixin api diff --git a/task-sdk/docs/mapped-tasks-vs-iterable-tasks.rst b/task-sdk/docs/mapped-tasks-vs-iterable-tasks.rst new file mode 100644 index 0000000000000..74d2515819619 --- /dev/null +++ b/task-sdk/docs/mapped-tasks-vs-iterable-tasks.rst @@ -0,0 +1,516 @@ + .. 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. + +.. _sdk-mapped-tasks-vs-iterable-tasks: + +Mapped tasks vs iterable tasks +============================== + +.. versionadded:: 3.4.0 + +Airflow provides two complementary ways to process collections of data: + +- **Mapped tasks** distribute work **across multiple workers**. + Each item becomes a separate Task Instance that can run on a different worker, + giving you horizontal scalability and per-item observability. + +- **Iterable Tasks (IT)** improves concurrency **within a single task**. + All items are processed inside one Task Instance on one worker, eliminating + scheduling overhead and — when combined with async operators — enabling true + I/O multiplexing through a shared event loop. + +In short: **mapping spreads load across workers; IT speeds up work within one worker.** + +While both approaches allow you to apply an operation over a collection, +they differ significantly in execution model, scheduler impact, and observability. +This page explains the trade-offs and when to use each. + +Real-World Motivation +--------------------- + +Consider a workflow that downloads ~17,000 XML files from an SFTP server and loads +them into a data warehouse. Community benchmarks compare a mapped operator with a loop +written by hand inside a single ``@task``, the pattern IT runs for you (none of these rows +uses ``iterate()`` itself): + +.. list-table:: + :header-rows: 1 + + * - Approach + - Execution Time + * - Mapped ``SFTPOperator`` + - 3 h 25 m + * - Sync ``@task`` with ``SFTPHook`` (sequential loop) + - 1 h 21 m + * - Async ``@task`` with ``SFTPHookAsync`` (concurrent loop) + - 8 m 29 s + * - Async ``@task`` with ``SFTPHookAsync`` and connection pooling + - 3 m 32 s + +The ~60× improvement comes from eliminating per-item scheduling overhead and +sharing a single event loop for concurrent I/O. This is the kind of workload +IT is built for: many small, I/O-bound operations processed within one task. + +Mapped tasks +------------ + +Task mapping allows you to expand a single task definition into multiple +Task Instances (TIs). + +For more details, see :ref:`task mapping `. + +Key characteristics: + +- Each item in the iterable creates a separate Task Instance. +- The scheduler is responsible for creating and managing all mapped tasks. +- Tasks can run in parallel across multiple worker slots. +- Fine-grained retry, logging, and observability per item. +- Well suited for workloads where each item should be independently scheduled and tracked. + +The following example fetches Pokémon data from a REST API. Each Pokémon becomes +a separate Task Instance, individually scheduled, retried, and visible in the UI: + +.. code-block:: python + + from datetime import datetime + + from airflow.providers.http.operators.http import HttpOperator + from airflow.sdk import DAG, task + + with DAG(dag_id="dtm-http-pokemon-example", start_date=datetime(2026, 1, 1)): + list_pokemon_task = HttpOperator( + task_id="list_pokemon", + http_conn_id="pokeapi", + method="GET", + endpoint="api/v2/pokemon?limit=100", + response_filter=lambda response: [ + pokemon["url"].replace("https://pokeapi.co/", "") for pokemon in response.json()["results"] + ], + log_response=False, + ) + + get_pokemon_task = HttpOperator.partial( + task_id="get_pokemon", + http_conn_id="pokeapi", + method="GET", + ).expand(endpoint=list_pokemon_task.output) + + list_pokemon_task >> get_pokemon_task + + +With 100 Pokémon the scheduler creates 100 Task Instances, each occupying +a worker slot. This is fine for small lists, but for thousands of items the +scheduler and database overhead becomes significant. + +Iterable Tasks (IT) +---------------------------- + +Iterable Tasks allows you to iterate over an iterable (typically an XCom result) +*within a single Task Instance*, applying an operator multiple times without creating +separate Task Instances. + +This means that iteration happens inside the task execution itself rather than at the +scheduler level. + +Key characteristics: + +- A single Task Instance processes all items in the iterable. +- No task expansion; the scheduler manages only one task. +- Lower scheduler overhead compared to mapping. +- Iterations share the same execution context (e.g., memory, event loop). +- Particularly well suited for async operators and high-throughput workloads. + +The same Pokémon fetching problem can be solved with IT. Here, a single Task +Instance processes all Pokémon concurrently using the sync +:class:`~airflow.providers.http.operators.http.HttpOperator`: + +.. code-block:: python + + from datetime import datetime + + from airflow.providers.http.operators.http import HttpOperator + from airflow.sdk import DAG, task + + with DAG(dag_id="it-http-pokemon-example", start_date=datetime(2026, 1, 1)): + list_pokemon_task = HttpOperator( + task_id="list_pokemon", + http_conn_id="pokeapi", + method="GET", + endpoint="api/v2/pokemon?limit=100", + response_filter=lambda response: [ + pokemon["url"].replace("https://pokeapi.co/", "") for pokemon in response.json()["results"] + ], + log_response=False, + ) + + get_pokemon_task = HttpOperator.partial( + task_id="get_pokemon", + http_conn_id="pokeapi", + method="GET", + ).iterate(endpoint=list_pokemon_task.output) + + list_pokemon_task >> get_pokemon_task + + +The scheduler only manages a single task. With a sync operator, iterations run in +a pool of up to ``task_concurrency`` threads, so blocking I/O such as these HTTP +requests overlaps up to that many at a time. Async operators scale further: a +coroutine waiting on I/O costs far less than a thread, so ``task_concurrency`` can +be set much higher. CPU-bound Python code speeds up with neither, since the +iterations share one process and its GIL. + +To **multiplex** many I/O-bound operations on one event loop, use an async task with +:class:`~airflow.providers.http.hooks.http.HttpAsyncHook`: + +.. code-block:: python + + from datetime import datetime + + from airflow.providers.http.hooks.http import HttpAsyncHook, HttpHook + from airflow.sdk import dag, task + + + @dag( + dag_id="it-async-http-pokemon-example", + start_date=datetime(2026, 1, 1), + ) + def it_async_http_pokemon_example(): + @task + def list_pokemon() -> list[str]: + response = HttpHook( + http_conn_id="pokeapi", + method="GET", + ).run( + endpoint="api/v2/pokemon?limit=100", + ) + + return [pokemon["url"].replace("https://pokeapi.co/", "") for pokemon in response.json()["results"]] + + @task( + retries=3, + task_concurrency=2, + show_return_value_in_logs=False, + ) + async def get_pokemon(url: str): + async with HttpAsyncHook( + http_conn_id="pokeapi", + method="GET", + ).session() as session: + response = await session.run(endpoint=url) + return await response.json() + + get_pokemon.iterate( + url=list_pokemon(), + ) + + + it_async_http_pokemon_example() + + +When ``iterate()`` is used with an async task, all iterations share the same +event loop, enabling true multiplexing of I/O-bound operations without any +manual concurrency management by the DAG author. For a handful of items the +difference is negligible, but for hundreds or thousands of items the +concurrent approach is dramatically faster — see the +:ref:`benchmarks above `. + +.. note:: + + ``multiple_outputs`` is ignored by ``iterate()``. Each iteration's return value is pushed + whole as ``return_value_`` and the task's own return value is the lazy sequence over + them, so a ``dict`` return annotation on the task does not fan its keys out into separate + XComs the way it does with ``expand()``. Every key an iteration pushes or stores carries its + index the same way: ``ti.xcom_push("foo", v)`` in iteration 2 lands under ``foo_2``, and so + does ``task_state_store.set("foo", v)``, so iterations never overwrite each other's values. + Reading them back follows the same rule: ``ti.xcom_pull(key="foo")`` in iteration 2, or with + its own ``task_ids``, reads ``foo_2``, as ``task_state_store.get("foo")`` does; a pull from + another task keeps its key. To read another iteration's value, use ``XCom.get_one`` + (``XCom.aget_one`` in an async operator) with the suffixed key, or read every value downstream + through the task's lazy sequence. + This holds in the iteration's own thread or coroutine. A ``threading.Thread`` the task starts, + or ``loop.run_in_executor()``, does not inherit it: ``get_current_context()`` there returns + the task's own context, whose ``ti`` and ``task_state_store`` add no index, so the iterations' + keys overwrite each other. Use ``context["ti"]`` passed to the task, or start helpers with + ``asyncio.to_thread()`` or ``contextvars.copy_context().run()``, which carry the iteration's + context over. + +.. warning:: + + Inputs are shared between iterations. All iterations run in one process, so a value handed + to several of them is the same object in each: every value passed through ``partial()``, and + with ``iterate(a=..., b=...)`` every element of ``a`` and of ``b``, which the cross product + combines more than once. With ``expand()`` each task instance runs in its own process and + gets its own copy. Treat inputs as read-only, or copy what the task changes in place. + +Why Iterable Tasks? +--------------------------- + +IT is designed to address limitations of task mapping in specific scenarios: + +- **Scheduler scalability**: + Mapping creates one Task Instance per item, which can put pressure on the scheduler + for very large datasets. IT avoids this by keeping execution within a single task. + +- **Async multiplexing**: + With Python-native async support in Airflow 3.2, IT allows multiple + operations to share the same event loop within a single Task Instance. + This enables efficient multiplexing of I/O-bound workloads. + +- **Lower overhead**: + No need to serialize, schedule, and track thousands of Task Instances. + +- **Triggerer and deferrable-operator bottleneck**: + Deferrable operators delegate async work to triggerers, which store yielded + events directly in the Airflow metadata database. Unlike workers, triggerers + cannot leverage a custom XCom backend to offload large payloads. This makes + triggerers a bottleneck for sustained high-load async execution or workloads + that return large results. Mapping deferrable operators + amplifies the problem further. IT sidesteps triggerers entirely — iterations + execute on workers, which scale more effectively and support custom XCom + backends. + + A custom XCom backend is not the whole story for IT, though: each item's result + is also written to its checkpoint in the task state store, so that a retry can + replay it instead of running the item again. Those checkpoints land in the + ``task_state_store`` table of the metadata database and stay there for the + store's retention (``[state_store] default_retention_days``, 30 days unless + configured) unless a ``[workers] state_store_backend`` is configured, in + which case the table only holds a reference to the payload. Iterating over + large results therefore needs both a custom XCom backend and a state store + backend; with only the first, the payloads move from the XCom table to the + task state store table. + + For more on deferred vs async trade-offs, see :doc:`deferred-vs-async-operators`. + +IT is especially useful for patterns such as: + +- API pagination +- Bulk HTTP or database calls +- High-throughput async workloads +- Streaming or lazily-evaluated XCom results + +Hooks as Building Blocks +^^^^^^^^^^^^^^^^^^^^^^^^ + +IT encourages a pattern where DAG authors call **hooks** directly from +``@task``-decorated functions rather than relying on operators. Operators are +wrappers around hooks and sometimes expose only a subset of the hook's +capabilities. By calling hooks directly, users gain full control over +concurrency, error handling, and batching. + +For example, instead of using ``HttpOperator`` in deferrable mode (which +delegates to the triggerer for a single request at a time), an async +``@task`` can call :class:`~airflow.providers.http.hooks.http.HttpAsyncHook` +directly to perform many concurrent requests. With IT, the framework +handles the iteration, concurrency, and event-loop management +automatically — the DAG author only writes the per-item logic and decides +which strategy it wants to use. + +This "hooks as building blocks" approach is especially powerful with async +hooks, where the shared event loop enables concurrent I/O without any +manual ``asyncio.gather`` or ``asyncio.Semaphore`` management. + +For more examples of calling async hooks directly from tasks, see +:doc:`deferred-vs-async-operators`. + +Callbacks +--------- + +With IT the callbacks of the wrapped operator run per item, against the item's own context, but +not all at the same moment: + +* ``on_success_callback`` and ``on_skipped_callback`` run once the item's checkpoint is written, so they + speak for work a retry will not run again; a checkpoint write that fails fires nothing, and the + attempt that runs the item again reports it then. +* ``on_failure_callback`` and ``on_retry_callback`` of a failed item wait until every item has run, + because whether the task is retried depends on all of them: an ``AirflowFailException`` in one + item fails the whole task without a retry. Once the task's fate is known, every failed item gets + the callback that matches it, one after another: ``on_retry_callback`` when the task is retried, + ``on_failure_callback`` when it is not. +* A failure that belongs to no item fires no callback: an error while resolving the input, before + any item exists, or items cancelled because the task's ``execution_timeout`` ran out (the item + the timeout struck is reported like the other failures), or items pulled into a free slot before + a kill and never started. The iterated task carries no task-level callbacks of its own. +* Listeners (``on_task_instance_running``, ``on_task_instance_success``, + ``on_task_instance_failed``) fire once, for the iterated task instance, when the runner reports + its state; an item is not a task instance and fires none. With ``.expand()`` they fire once per + mapped task instance. + +Comparison +---------- + +.. list-table:: + :header-rows: 1 + + * - Aspect + - Mapped tasks + - Iterable Tasks (IT) + * - Task Instances + - One per item + - Single Task Instance + * - Scheduler load + - High for large iterables + - Minimal + * - Execution model + - Distributed across workers + - In-process iteration + * - Concurrency + - Parallel tasks + - Sync or async within one task + * - Async support + - Limited (per task) + - Strong (shared event loop, multiplexing) + * - Retry behavior + - Per item + - Whole task retries, but checkpointed items are skipped + * - Skipped items + - A skipped task instance pushes no XCom. Downstream tasks with ``all_success`` are skipped; + with ``none_failed`` they run over the other items + - The same: a skipped iteration is left out of the result, downstream tasks with + ``all_success`` are skipped, and with ``none_failed`` they run over the other items + * - Items that return ``None`` + - The task instance pushes no XCom and is not counted: a downstream ``.expand()`` over the + output runs over the values that exist, and positions shift + - The iteration pushes no XCom but keeps its position: the result reads ``None`` there, and a + downstream ``.expand()`` over it runs over that ``None`` + * - Asset events + - One event per mapped task instance that emits to an asset + - One event per asset for the whole task instance: items emitting to the same asset are + merged into it, their ``extra`` with the last item to finish winning, partitions and alias + events accumulated. A per-file ``Metadata`` pattern from ``.expand()`` produces one event here + * - Empty input + - The mapped task is skipped, and so are downstream tasks with ``all_success`` + - The same: the task is skipped and pushes no result + * - Observability + - Per item in UI + - Aggregated in a single task + * - Triggerer dependency + - Deferrable mapped tasks rely on triggerers + - No triggerers involved + * - Deferral / reschedule + - Supported (each item has its own task instance) + - Not supported (raises a non-retryable failure) + * - XCom backend + - Workers support custom XCom backends + - Workers support custom XCom backends (triggerers do not); each item's result is also + checkpointed in the task state store, so large results need a ``state_store_backend`` too + * - Use case + - Independent, trackable units of work + - High-throughput or streaming workloads + +The following table illustrates these differences using the Pokémon example from above: + +.. list-table:: + :header-rows: 1 + + * - Pattern + - Task Instances + - Work Per Task + * - ``get_pokemon.expand(url=urls)`` + - 100 + - 1 Pokémon + * - ``get_pokemon.iterate(url=urls)`` + - 1 + - 100 Pokémon + +When to Use Mapped Tasks +------------------------ + +Prefer mapped tasks when: + +- Each item must be independently tracked in the UI. +- You need fine-grained retries per item. +- Tasks are long-running or resource-intensive. +- Work should be distributed across multiple workers. +- Scheduling decisions should be made per item. +- You need deferrable operators or reschedule-mode sensors — these work natively with mapped tasks, + since each mapped item has its own task instance to defer or reschedule. + +When to Use Iterable Tasks +----------------------------------- + +Prefer IT when: + +- You are processing large numbers of small items. +- Scheduler overhead becomes a concern. +- You are using async operators and want to leverage a shared event loop. +- Workloads are I/O-bound and benefit from multiplexing. +- Fine-grained observability per item is not required. + +When **not** to use IT +----------------------- + +Avoid Iterable Tasks when: + +- Each item represents a long-running or CPU-bound computation: the iterations share one + process, so the GIL keeps CPU-bound Python code from running in parallel. +- You require detailed visibility per item in the Airflow UI. +- Work must be distributed across multiple worker nodes. +- Sub-tasks need to defer (deferrable operators) or reschedule (reschedule-mode sensors) — a + sub-task index has no task instance of its own to defer or reschedule against, so either raises + a non-retryable failure instead of pausing. +- The operator finds its remote work by the task instance's identity. Every item runs as the one + iterated task instance, with its ``dag_id``, ``task_id``, ``run_id`` and ``map_index``: a + ``KubernetesPodOperator`` that reattaches (``durable``, or ``reattach_on_restart`` before it) + looks a running pod up by those labels and adopts a sibling item's pod, or fails on finding + several; an ``EcsRunTaskOperator`` with ``reattach=True`` builds its ``startedBy`` from the same + fields. Where the operator takes labels of its own, one that names the item tells the jobs apart: + ``labels={"index": "{{ ti.index }}"}`` on the ``KubernetesPodOperator``, as every template of an + item renders against the item's own ``ti``. Otherwise turn reattachment off, or map the task with + ``.expand()``. + +.. tip:: + + IT is a **third execution option** alongside task mapping and + deferrable operators. It is not intended as a replacement for either. + Triggerers remain the right choice for long-running polling or waiting tasks + (e.g., monitoring a remote job or waiting for a Kubernetes pod to complete). + +Relationship with Async Operators +---------------------------------- + +IT complements async operators introduced in Airflow 3.2 and is the natural next step in that +evolution: async operators make a single I/O call non-blocking, while IT applies that same +non-blocking call repeatedly across a dataset within one task. + +- Async operators allow concurrent I/O within a single task. +- IT allows you to *apply an operator repeatedly* over a dataset within that same task. + +Together, they enable patterns such as: + +- Efficient API pagination +- Concurrent request batching +- Streaming data processing + +Unlike task mapping, where each mapped task runs in its own execution context, +IT allows all iterations to share the same event loop, enabling true multiplexing. + +Because IT executes on workers rather than triggerers, it also benefits from the +full worker environment: custom XCom backends, Edge Worker support, and the +scalability of execution frameworks such as Celery. + +For more details on async execution, see :doc:`deferred-vs-async-operators`. + +Future Outlook +-------------- + +As Python's async ecosystem evolves, IT tasks will benefit from improved +introspection and tooling. For example, Python 3.14 introduces new +`asyncio introspection capabilities `_ +that could eventually enable structured progress reporting in the Airflow UI +for IT tasks — providing per-item visibility without the overhead of per-item +task instances. diff --git a/task-sdk/src/airflow/sdk/bases/decorator.py b/task-sdk/src/airflow/sdk/bases/decorator.py index c4d8ad7b5a6ed..9b4bc3914f48b 100644 --- a/task-sdk/src/airflow/sdk/bases/decorator.py +++ b/task-sdk/src/airflow/sdk/bases/decorator.py @@ -41,6 +41,7 @@ from airflow.sdk.definitions._internal.decorators import remove_task_decorator from airflow.sdk.definitions._internal.expandinput import ( EXPAND_INPUT_EMPTY, + DecoratedExpandInput, DictOfListsExpandInput, ListOfDictsExpandInput, is_mappable, @@ -49,6 +50,7 @@ from airflow.sdk.definitions.asset import Asset from airflow.sdk.definitions.context import KNOWN_CONTEXT_KEYS from airflow.sdk.definitions.mappedoperator import ( + TASK_CONCURRENCY_REJECTED, MappedOperator, ensure_xcomarg_return_value, prevent_duplicates, @@ -96,10 +98,10 @@ def _validate_arg_names(self, func: ValidationSource, kwargs: dict[str, Any]) -> kwargs_left = kwargs.copy() for arg_name in self._mappable_function_argument_names: value = kwargs_left.pop(arg_name, NOTSET) - if func == "expand" and value is not NOTSET and not is_mappable(value): + if func in ("expand", "iterate") and value is not NOTSET and not is_mappable(value): tname = type(value).__name__ raise ValueError( - f"expand() got an unexpected type {tname!r} for keyword argument {arg_name!r}" + f"{func}() got an unexpected type {tname!r} for keyword argument {arg_name!r}" ) if len(kwargs_left) == 1: raise TypeError(f"{func}() got an unexpected keyword argument {next(iter(kwargs_left))!r}") @@ -590,6 +592,12 @@ def expand(self, **map_kwargs: OperatorExpandArgument) -> XComArg: ) if not map_kwargs: raise TypeError("no arguments to expand against") + # task_concurrency only has meaning for Iterable Tasks (as the sub-task thread + # count consumed by IterableOperator via .iterate()/.iterate_kwargs()). + # A plain .expand() never reaches that code path, so reject it here rather than silently + # accepting a dead value. + if "task_concurrency" in self.kwargs: + raise TypeError(TASK_CONCURRENCY_REJECTED) self._validate_arg_names("expand", map_kwargs) prevent_duplicates(self.kwargs, map_kwargs, fail_reason="mapping already partial") # Since the input is already checked at parse time, we can set strict @@ -598,7 +606,7 @@ def expand(self, **map_kwargs: OperatorExpandArgument) -> XComArg: if "trigger_rule" in self.kwargs: raise ValueError("Trigger rule not configurable for teardown tasks.") self.kwargs.update(trigger_rule=TriggerRule.ALL_DONE_SETUP_SUCCESS) - return self._expand(DictOfListsExpandInput(map_kwargs), strict=False) + return XComArg(operator=self._expand(DictOfListsExpandInput(map_kwargs), strict=False)) def expand_kwargs(self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = True) -> XComArg: if ( @@ -622,9 +630,18 @@ def expand_kwargs(self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") elif not isinstance(kwargs, XComArg): raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") - return self._expand(ListOfDictsExpandInput(kwargs), strict=strict) + # See the comment in expand() above: task_concurrency has no meaning outside iterate(). + if "task_concurrency" in self.kwargs: + raise TypeError(TASK_CONCURRENCY_REJECTED) + return XComArg(operator=self._expand(ListOfDictsExpandInput(kwargs), strict=strict)) - def _expand(self, expand_input: ExpandInput, *, strict: bool) -> XComArg: + def _expand( + self, + expand_input: ExpandInput, + *, + strict: bool, + register_with_dag: bool = True, + ) -> DecoratedMappedOperator: ensure_xcomarg_return_value(expand_input.value) task_kwargs = self.kwargs.copy() @@ -692,7 +709,7 @@ def _expand(self, expand_input: ExpandInput, *, strict: bool) -> XComArg: except AttributeError: operator_name = self.operator_class.__name__ - operator = _MappedOperator( + return _MappedOperator( operator_class=self.operator_class, expand_input=EXPAND_INPUT_EMPTY, # Don't use this; mapped values go to op_kwargs_expand_input. partial_kwargs=partial_kwargs, @@ -726,8 +743,69 @@ def _expand(self, expand_input: ExpandInput, *, strict: bool) -> XComArg: start_trigger_args=self.operator_class.start_trigger_args, start_from_trigger=self.operator_class.start_from_trigger, returns_dag_result=self.returns_dag_result, + register_with_dag=register_with_dag, + ) + + def iterate(self, **map_kwargs: OperatorExpandArgument) -> XComArg: + """ + Iterate the task over ``map_kwargs`` inside a single task instance. + + The counterpart of :meth:`expand` for Iterable Tasks: the same inputs, but processed by one + :class:`~airflow.sdk.definitions.iterableoperator.IterableOperator` instead of one task + instance per item. + """ + if self.kwargs.get("trigger_rule") == TriggerRule.ALWAYS and any( + [isinstance(expanded, XComArg) for expanded in map_kwargs.values()] + ): + raise ValueError( + "Task-generated iterating within a task using 'iterate' is not allowed with trigger rule 'always'." + ) + if not map_kwargs: + raise TypeError("no arguments to iterate against") + self._validate_arg_names("iterate", map_kwargs) + prevent_duplicates(self.kwargs, map_kwargs, fail_reason="mapping already partial") + # Since the input is already checked at parse time, we can set strict + # to False to skip the checks on execution. + if self.is_teardown: + if "trigger_rule" in self.kwargs: + raise ValueError("Trigger rule not configurable for teardown tasks.") + self.kwargs.update(trigger_rule=TriggerRule.ALL_DONE_SETUP_SUCCESS) + return self._iterate(DictOfListsExpandInput(map_kwargs), strict=False) + + def iterate_kwargs(self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = True) -> XComArg: + """Iterate the task over a list of dicts or an XComArg; see :meth:`iterate`.""" + if ( + self.kwargs.get("trigger_rule") == TriggerRule.ALWAYS + and not isinstance(kwargs, XComArg) + and any( + [ + isinstance(v, XComArg) + for kwarg in kwargs + if not isinstance(kwarg, XComArg) + for v in kwarg.values() + ] + ) + ): + raise ValueError( + "Task-generated iterating within a task using 'iterate_kwargs' is not allowed with trigger rule 'always'." + ) + if isinstance(kwargs, Sequence): + for item in kwargs: + if not isinstance(item, (XComArg, Mapping)): + raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") + elif not isinstance(kwargs, XComArg): + raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") + return self._iterate(ListOfDictsExpandInput(kwargs), strict=strict) + + def _iterate(self, expand_input: ExpandInput, *, strict: bool) -> XComArg: + from airflow.sdk.definitions.iterableoperator import IterableOperator + + # The DecoratedMappedOperator only drives the iteration in memory: it is never registered + # with the DAG, the IterableOperator is the single real task. + operator = self._expand(expand_input, strict=strict, register_with_dag=False) + return XComArg( + operator=IterableOperator(operator=operator, expand_input=DecoratedExpandInput(expand_input)) ) - return XComArg(operator=operator) def partial(self, **kwargs: Any) -> _TaskDecorator[FParams, FReturn, OperatorSubclass]: self._validate_arg_names("partial", kwargs) @@ -786,8 +864,8 @@ class Task(Protocol, Generic[FParams, FReturn]): An instance of this type inherits the call signature of the decorated function wrapped in it (not *exactly* since it actually returns an XComArg, - but there's no way to express that right now), and provides two additional - methods for task-mapping. + but there's no way to express that right now), and provides the methods for + task-mapping and task iteration. This type is implemented by ``_TaskDecorator`` at runtime. """ @@ -805,6 +883,10 @@ def expand(self, **kwargs: OperatorExpandArgument) -> XComArg: ... def expand_kwargs(self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = True) -> XComArg: ... + def iterate(self, **kwargs: OperatorExpandArgument) -> XComArg: ... + + def iterate_kwargs(self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = True) -> XComArg: ... + def override(self, **kwargs: Any) -> Task[FParams, FReturn]: ... diff --git a/task-sdk/src/airflow/sdk/bases/operator.py b/task-sdk/src/airflow/sdk/bases/operator.py index 536b4545d17c9..9bc423a7f7f4d 100644 --- a/task-sdk/src/airflow/sdk/bases/operator.py +++ b/task-sdk/src/airflow/sdk/bases/operator.py @@ -63,7 +63,11 @@ from airflow.sdk.definitions._internal.setup_teardown import SetupTeardownContext from airflow.sdk.definitions._internal.types import NOTSET, validate_instance_args from airflow.sdk.definitions.edges import EdgeModifier -from airflow.sdk.definitions.mappedoperator import OperatorPartial, validate_mapping_kwargs +from airflow.sdk.definitions.mappedoperator import ( + TASK_CONCURRENCY_REJECTED, + OperatorPartial, + validate_mapping_kwargs, +) from airflow.sdk.definitions.param import ParamsDict from airflow.sdk.exceptions import RemovedInAirflow4Warning @@ -325,6 +329,7 @@ def partial( map_index_template: str | None = ..., max_active_tis_per_dag: int | None = ..., max_active_tis_per_dagrun: int | None = ..., + task_concurrency: int | None = ..., on_execute_callback: None | TaskStateChangeCallback | list[TaskStateChangeCallback] = ..., on_failure_callback: None | TaskStateChangeCallback | list[TaskStateChangeCallback] = ..., on_success_callback: None | TaskStateChangeCallback | list[TaskStateChangeCallback] = ..., @@ -397,8 +402,6 @@ def partial( ) # Post-process arguments. Should be kept in sync with _TaskDecorator.expand(). - if "task_concurrency" in kwargs: # Reject deprecated option. - raise TypeError("unexpected argument: task_concurrency") if start_date := partial_kwargs.get("start_date", None): partial_kwargs["start_date"] = timezone.convert_to_utc(start_date) if end_date := partial_kwargs.get("end_date", None): @@ -841,6 +844,13 @@ class derived from this one results in the creation of a task object, key in the returned dictionary result. If False and do_xcom_push is True, pushes a single XCom. :param task_group: The TaskGroup to which the task should belong. This is typically provided when not using a TaskGroup as a context manager. + :param task_concurrency: How many iterations run at once when the operator is iterated with + ``.iterate()`` or ``.iterate_kwargs()``; rejected everywhere else. Defaults to + ``os.cpu_count()``, the CPUs of the machine whatever CPU limit the worker's container has. + Not the Airflow 2 option of the same name, which limited concurrent task instances and is + now ``max_active_tis_per_dag``. Only read when passed to the task itself (``.partial()`` or + ``@task``); ``default_args`` do not set it, as they did not before, so a DAG that still + carries the Airflow 2 option there is unaffected. :param doc: Add documentation or notes to your Task objects that is visible in Task Instance details View in the Webserver :param doc_md: Add documentation (in Markdown format) or notes to your Task objects @@ -1120,6 +1130,14 @@ def __init__( super().__init__() self.task_group = task_group + # task_concurrency only has meaning for Iterable Tasks (as the sub-task thread + # count, see IterableOperator.max_workers). A directly instantiated operator can never + # reach that code path, so reject it here rather than silently accepting a dead value. + # It is deliberately not a parameter of this method: apply_defaults copies every parameter + # found in default_args into the call, which would reject it for every task of a DAG + # carrying the Airflow 2 option in its default_args. + if kwargs.pop("task_concurrency", None) is not None: + raise TypeError(TASK_CONCURRENCY_REJECTED) kwargs.pop("_airflow_mapped_validation_only", None) if kwargs: diff --git a/task-sdk/src/airflow/sdk/bases/xcom.py b/task-sdk/src/airflow/sdk/bases/xcom.py index 50e54c458b087..c2b5f38565b95 100644 --- a/task-sdk/src/airflow/sdk/bases/xcom.py +++ b/task-sdk/src/airflow/sdk/bases/xcom.py @@ -18,6 +18,7 @@ from __future__ import annotations import collections +from collections.abc import AsyncIterator, Iterator, Sequence from typing import Any, Protocol import structlog @@ -564,3 +565,166 @@ def delete( ), ) cls.purge(xcom_result) + + +def _normalize_index(index: int, length: int) -> int: + """Map a sequence index, negative ones included, onto a position in ``[0, length)``.""" + if index < 0: + index += length + if not (0 <= index < length): + raise IndexError(index) + return index + + +class XComIterable(Sequence): + """ + An iterable that lazily fetches XCom values one by one instead of loading all at once. + + This is a read-only :class:`collections.abc.Sequence` over the ``return_value_`` XComs an + iterated task pushed, one per index: the values are written by the producing task's runner as + each sub-task finishes (see ``IterableOperator.axcom_push``), and the iterable only ever reads + them. Nothing on this class mutates the underlying XComs. + + Indexing follows the usual sequence rules, negative indices included: ``result[-1]`` is the last + value. Iterations that were skipped pushed nothing and are left out, as the XComs of skipped + mapped task instances are: ``length`` counts every input item, ``skipped`` lists the indices + that were skipped, and positions in the sequence run over the others only. An iteration that + returned ``None`` pushed nothing either but keeps its position: reading it gives ``None``, + where ``.expand()`` would not count such a mapped task instance at all. + + Every element is a remote fetch, so random access costs one XCom read per element, and so do + iterating and slicing: N elements are N requests. Reading several in one request needs an + Execution API endpoint that takes a list of keys, which does not exist yet. + """ + + def __init__( + self, + task_id: str, + dag_id: str, + run_id: str, + map_index: int | None = None, + length: int | None = None, + skipped: Sequence[int] = (), + ): + self.task_id = task_id + self.dag_id = dag_id + self.run_id = run_id + self.map_index = map_index + self.length = length or 0 + self.skipped: list[int] = sorted(skipped) + + def _index_of(self, position: int) -> int: + """Map ``position`` in the sequence to its input index, stepping over the skipped indices.""" + index = _normalize_index(position, len(self)) + for skipped_index in self.skipped: + if skipped_index > index: + break + index += 1 + return index + + def __iter__(self) -> Iterator[Any]: + return _XComIterator(self) + + def __len__(self) -> int: + return self.length - len(self.skipped) + + async def alen(self) -> int: + """Async twin of ``len(self)``, for readers on the event loop that take a length before each read.""" + return len(self) + + def __getitem__(self, key: int | slice) -> Any | Sequence[Any]: + """Allow direct indexing so this works like a sequence.""" + from airflow.sdk.execution_time.xcom import XCom + + if isinstance(key, slice): + # TODO: This issues one XCom.get_one call per element — N round-trips for a full slice. + # XComIterable stores results under distinct keys (return_value_0, return_value_1, …) + # with the same map_index, so the existing GetXComSequenceSlice endpoint (which ranges + # over map_index for a single key) cannot be reused. A new POST endpoint that accepts + # a list of keys and returns values in a single query is needed; once that lands, replace + # this loop with a single batched fetch. + start, stop, step = key.indices(len(self)) + return [self[i] for i in range(start, stop, step)] + + return XCom.get_one( + key=f"{BaseXCom.XCOM_RETURN_KEY}_{self._index_of(key)}", + dag_id=self.dag_id, + task_id=self.task_id, + run_id=self.run_id, + map_index=self.map_index, + ) + + async def aget(self, index: int) -> Any: + """ + Async counterpart of ``self[index]``: fetch one value through ``XCom.aget_one``. + + Use it, or ``async for``, from code running on an event loop that has other SDK calls in + flight (an iterated task consuming this iterable as its input): a synchronous read there + would block the loop thread on the supervisor channel and deadlock with them. + """ + from airflow.sdk.execution_time.xcom import XCom + + return await XCom.aget_one( + key=f"{BaseXCom.XCOM_RETURN_KEY}_{self._index_of(index)}", + dag_id=self.dag_id, + task_id=self.task_id, + run_id=self.run_id, + map_index=self.map_index, + ) + + def __aiter__(self) -> AsyncIterator[Any]: + return _AsyncXComIterator(self) + + def serialize(self) -> dict: + """Ensure the object is JSON serializable.""" + return { + "task_id": self.task_id, + "dag_id": self.dag_id, + "run_id": self.run_id, + "map_index": self.map_index, + "length": self.length, + "skipped": self.skipped, + } + + @classmethod + def deserialize(cls, data: dict, version: int): + """Ensure the object is JSON deserializable.""" + return cls(**data) + + +class _AsyncXComIterator: + """Async iterator for XComIterable, one ``aget`` per position in order.""" + + def __init__(self, iterable: XComIterable): + self._iterable = iterable + self._index = 0 + + def __aiter__(self): + return self + + async def __anext__(self): + if self._index >= await self._iterable.alen(): + raise StopAsyncIteration + + value = await self._iterable.aget(self._index) + self._index += 1 + return value + + +class _XComIterator: + """Iterator for XComIterable.""" + + def __init__(self, iterable: XComIterable): + self._iterable = iterable + self._index = 0 + + def __iter__(self): + return self + + def __next__(self): + if self._index >= len(self._iterable): + raise StopIteration + + value = self._iterable[self._index] + self._index += 1 + return value diff --git a/task-sdk/src/airflow/sdk/definitions/_internal/contextmanager.py b/task-sdk/src/airflow/sdk/definitions/_internal/contextmanager.py index f0fbdc306e6f3..dbf6998852fc0 100644 --- a/task-sdk/src/airflow/sdk/definitions/_internal/contextmanager.py +++ b/task-sdk/src/airflow/sdk/definitions/_internal/contextmanager.py @@ -19,6 +19,7 @@ import sys from collections import deque +from contextvars import ContextVar from types import ModuleType from typing import TYPE_CHECKING, Any, Generic, TypeVar @@ -38,8 +39,17 @@ # the `get_current_context` function. _CURRENT_CONTEXT: list[Context] = [] +# The contexts of the iterations of an iterated task, pushed on top of the task's own context. +# Iterations run concurrently in one process, in worker threads and asyncio tasks, so each one's +# context lives in a ContextVar and is seen by its own execution only. The task's context stays in +# the module-level list above, where any thread of the process finds it, including one the task +# started itself, whose ContextVars start out empty. +_INDEXED_CONTEXT: ContextVar[tuple[Context, ...]] = ContextVar("_indexed_context", default=()) + def _get_current_context() -> Context: + if indexed := _INDEXED_CONTEXT.get(): + return indexed[-1] if not _CURRENT_CONTEXT: raise RuntimeError( "Current context was requested but no context was found! Are you running within an Airflow task?" diff --git a/task-sdk/src/airflow/sdk/definitions/_internal/expandinput.py b/task-sdk/src/airflow/sdk/definitions/_internal/expandinput.py index b6ffbd2214253..7fe210118b30f 100644 --- a/task-sdk/src/airflow/sdk/definitions/_internal/expandinput.py +++ b/task-sdk/src/airflow/sdk/definitions/_internal/expandinput.py @@ -17,21 +17,22 @@ # under the License. from __future__ import annotations -from collections.abc import Iterable, Mapping, Sequence, Sized -from typing import TYPE_CHECKING, Any, ClassVar, Union +import asyncio +import math +from abc import ABC, abstractmethod +from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence, Sized +from typing import TYPE_CHECKING, Any, ClassVar, NamedTuple, Union import attrs from airflow.sdk.definitions._internal.mixins import ResolveMixin +from airflow.sdk.definitions.xcom_arg import XComArg if TYPE_CHECKING: from typing import TypeGuard - from airflow.sdk.definitions.xcom_arg import XComArg from airflow.sdk.types import Operator -ExpandInput = Union["DictOfListsExpandInput", "ListOfDictsExpandInput"] - # Each keyword argument to expand() can be an XComArg, sequence, or dict (not # any mapping since we need the value to be ordered). OperatorExpandArgument = Union["MappedArgument", "XComArg", Sequence, dict[str, Any]] @@ -41,6 +42,96 @@ OperatorExpandKwargsArgument = Union["XComArg", Sequence[Union["XComArg", Mapping[str, Any]]]] +class Resolved(NamedTuple): + """ + An expand input resolved for iteration: its item count and an async read by index. + + The ``.iterate()`` counterpart of ``ExpandInput.resolve``, which hands one mapped task + instance the item at its ``map_index``: every source is pulled once, and ``aget(index)`` then + picks the item an index maps to the way ``resolve`` does for ``.expand()`` (the cross product + of ``iterate(**kwargs)``, the mapping at that position for ``iterate_kwargs``). Reads run on + the task's event loop next to the sub-tasks' own SDK calls, so nothing here may block the loop + thread on the supervisor channel (see ``AsyncAwareExecutor.imap_unordered``). + """ + + length: int + aget: Callable[[int], Awaitable[Mapping[str, Any]]] + + +class Source: + """ + One resolved expand argument, read by index without a blocking supervisor call on the loop. + + A Mapping is read as its ``(key, value)`` pairs and a scalar as a one-item sequence, matching + ``DictOfListsExpandInput._expand_mapped_field``. A value with async accessors (an + ``XComIterable``, a mapped upstream's ``LazyXComSequence``) is read through them. Anything + else may block in ``__getitem__`` (a ``.map()`` result over a lazy sequence, for instance), so + it is read in a worker thread. + """ + + @classmethod + async def from_argument(cls, argument: Any, context: Mapping[str, Any]) -> Source: + """Build a source from an expand argument: pull it if it is an XComArg, then make it a sequence.""" + if isinstance(argument, XComArg): + argument = await argument.aresolve(context) + # What .expand() refuses when the upstream pushes (_push_xcom_if_needed raises for a + # mapped dependant) is refused here: an IterableOperator is not a MappedOperator, so + # iter_mapped_dependants never finds it and that check does not fire for it. + from airflow.sdk.definitions.mappedoperator import is_mappable_value + from airflow.sdk.exceptions import UnmappableXComTypePushed, XComForMappingNotPushed + + if argument is None: + raise XComForMappingNotPushed() + if not is_mappable_value(argument): + raise UnmappableXComTypePushed(argument) + if isinstance(argument, Mapping): + return cls(list(argument.items())) + if isinstance(argument, (str, bytes)) or not isinstance(argument, Iterable): + # A literal is refused at parse time (validate_mapping_kwargs, _validate_arg_names), as + # .expand() refuses it, and an upstream value above; nothing else reaches this. + raise TypeError(f"cannot iterate over a {type(argument).__name__!r} argument") + if not isinstance(argument, Sequence): + return cls(list(argument)) + return cls(argument) + + def __init__(self, value: Sequence[Any]) -> None: + self.value = value + + async def alen(self) -> int: + """Item count: the value's own async count when it has one, else ``len()`` in a worker thread.""" + if hasattr(self.value, "alen"): + return await self.value.alen() + return await asyncio.to_thread(len, self.value) + + async def aget(self, index: int) -> Any: + """Item at ``index``; called once per item, so in-memory containers skip the thread hop.""" + value = self.value + if isinstance(value, (list, tuple, range)): + return value[index] + if hasattr(value, "aget"): + return await value.aget(index) + return await asyncio.to_thread(value.__getitem__, index) + + +def index_for_each_field(map_index: int, lengths: Mapping[str, int]) -> dict[str, int]: + """ + Split a cross-product position into one index per expand argument. + + The arguments are combined as ``itertools.product`` combines them in the order they were + given, so the last one varies fastest: position 3 of ``a=[1, 2], b=[10, 20]`` is ``a[1], b[1]``. + Shared by ``.expand()`` (``_expand_mapped_field`` picks its task instance's ``map_index``) and + ``.iterate()`` (``aresolve`` reads every position), so both hand a sub-task the same item. + """ + indices: dict[str, int] = {} + for key in reversed(list(lengths)): + length = lengths[key] + if length < 1: + raise RuntimeError(f"cannot expand field mapped to length {length!r}") + indices[key] = map_index % length + map_index //= length + return {key: indices[key] for key in lengths} + + class _NotFullyPopulated(RuntimeError): """ Raise when an expand input cannot be resolved due to incomplete metadata. @@ -79,6 +170,56 @@ def _needs_run_time_resolution(v: OperatorExpandArgument) -> TypeGuard[MappedArg return isinstance(v, (MappedArgument, XComArg)) +@attrs.define(slots=False) +class ExpandInput(ABC, ResolveMixin): + EXPAND_INPUT_TYPE: ClassVar[str] + + @property + @abstractmethod + def value(self) -> Any: + """The value of the expand input.""" + ... + + async def aresolve(self, context: Mapping[str, Any]) -> Resolved: + """ + Resolve every index of the input for an iterated task; see :class:`Resolved`. + + Implementations must not make a blocking supervisor call on the loop thread: XComArg + sources are pulled with ``XComArg.aresolve`` and read through :class:`Source`. + """ + raise NotImplementedError() + + def resolve(self, context: Mapping[str, Any]) -> Any: + raise NotImplementedError() + + +@attrs.define(slots=False) +class DecoratedExpandInput(ExpandInput): + """The expand input of a decorated task, whose items arrive as ``op_kwargs``.""" + + EXPAND_INPUT_TYPE: ClassVar[str] = "decorated" + + delegate: ExpandInput + + @property + def value(self) -> Any: + return self.delegate.value + + def iter_references(self) -> Iterable[tuple[Operator, str]]: + return self.delegate.iter_references() + + async def aresolve(self, context: Mapping[str, Any]) -> Resolved: + length, aget = await self.delegate.aresolve(context) + + async def aget_op_kwargs(index: int) -> Mapping[str, Any]: + return {"op_kwargs": await aget(index)} + + return Resolved(length, aget_op_kwargs) + + def resolve(self, context: Mapping[str, Any]) -> tuple[Mapping[str, Any], set[int]]: + return self.delegate.resolve(context) + + @attrs.define(kw_only=True) class MappedArgument(ResolveMixin): """ @@ -107,7 +248,7 @@ def resolve(self, context: Mapping[str, Any]) -> Any: @attrs.define() -class DictOfListsExpandInput(ResolveMixin): +class DictOfListsExpandInput(ExpandInput): """ Storage type of a mapped operator's mapped kwargs. @@ -154,20 +295,8 @@ def _get_length(k: str, v: OperatorExpandArgument) -> int | None: return map_lengths def _expand_mapped_field(self, key: str, value: Any, map_index: int, all_lengths: dict[str, int]) -> Any: - def _find_index_for_this_field(index: int) -> int: - # Need to use the original user input to retain argument order. - for mapped_key in reversed(self.value): - mapped_length = all_lengths[mapped_key] - if mapped_length < 1: - raise RuntimeError(f"cannot expand field mapped to length {mapped_length!r}") - if mapped_key == key: - return index % mapped_length - index //= mapped_length - return -1 - - found_index = _find_index_for_this_field(map_index) - if found_index < 0: - return value + # Use the original user input to retain argument order. + found_index = index_for_each_field(map_index, {k: all_lengths[k] for k in self.value})[key] if isinstance(value, Sequence): return value[found_index] if not isinstance(value, dict): @@ -184,6 +313,16 @@ def iter_references(self) -> Iterable[tuple[Operator, str]]: if isinstance(x, XComArg): yield from x.iter_references() + async def aresolve(self, context: Mapping[str, Any]) -> Resolved: + sources = {key: await Source.from_argument(value, context) for key, value in self.value.items()} + lengths = {key: await source.alen() for key, source in sources.items()} + + async def aget(index: int) -> Mapping[str, Any]: + positions = index_for_each_field(index, lengths) + return {key: await source.aget(positions[key]) for key, source in sources.items()} + + return Resolved(math.prod(lengths.values()), aget) + def resolve(self, context: Mapping[str, Any]) -> tuple[Mapping[str, Any], set[int]]: map_index: int | None = context["ti"].map_index if map_index is None or map_index < 0: @@ -217,7 +356,7 @@ def _describe_type(value: Any) -> str: @attrs.define() -class ListOfDictsExpandInput(ResolveMixin): +class ListOfDictsExpandInput(ExpandInput): """ Storage type of a mapped operator's mapped kwargs. @@ -238,12 +377,29 @@ def iter_references(self) -> Iterable[tuple[Operator, str]]: if isinstance(x, XComArg): yield from x.iter_references() + async def aresolve(self, context: Mapping[str, Any]) -> Resolved: + if isinstance(self.value, XComArg): + source = await Source.from_argument(self.value, context) + else: + source = Source( + [await item.aresolve(context) if isinstance(item, XComArg) else item for item in self.value] + ) + + async def aget(index: int) -> Mapping[str, Any]: + mapping = await source.aget(index) + if not isinstance(mapping, Mapping): + raise ValueError( + f"iterate_kwargs() expects a list[dict], not list[{_describe_type(mapping)}]" + ) + return mapping + + return Resolved(await source.alen(), aget) + def resolve(self, context: Mapping[str, Any]) -> tuple[Mapping[str, Any], set[int]]: map_index = context["ti"].map_index - if map_index < 0: + if map_index is None or map_index < 0: raise RuntimeError("can't resolve task-mapping argument without expanding") - mapping: Any = None if isinstance(self.value, Sized): mapping = self.value[map_index] if not isinstance(mapping, Mapping): diff --git a/task-sdk/src/airflow/sdk/definitions/context.py b/task-sdk/src/airflow/sdk/definitions/context.py index 98c8051a24761..c3d1733020747 100644 --- a/task-sdk/src/airflow/sdk/definitions/context.py +++ b/task-sdk/src/airflow/sdk/definitions/context.py @@ -93,6 +93,74 @@ class Context(TypedDict, total=False): KNOWN_CONTEXT_KEYS: set[str] = set(Context.__annotations__.keys()) +def clone_context(context: Context) -> Context: + """ + Create a safe, per-task copy of an execution ``Context`` for concurrent execution. + + The execution context is a mutable mapping that contains many nested + structures (``params``, ``templates_dict``, ``outlet_events``, ``dag_run``, + etc.). When running the same logical task concurrently (for example when + the ``IterableOperator`` spawns multiple indexed task instances that run in + parallel using threads, processes or asyncio tasks), those mutable objects + could be mutated by one indexed runtime and unintentionally observed by + another. That leads to subtle race conditions, corrupted state, and + flakiness in task execution. + + ``clone_context`` returns a new :class:`Context` mapping where the top-level + mapping is copied and specific mutable sub-objects that are commonly + mutated during execution are deep-copied or shallow-copied as appropriate: + + - ``params`` and ``templates_dict`` are deep-copied because they are + dictionaries that users and operators commonly mutate. + - ``inlets`` and ``outlets`` are converted to new lists (shallow copy) + because the sequence identity must be isolated but the elements are + typically read-only accessor objects. + - ``dag_run`` is intentionally **not** copied. It is treated as read-only + runtime metadata (e.g. ``run_id``, ``logical_date``, ``conf``) that all + concurrent sub-tasks legitimately share; copying it would only add + overhead without providing meaningful isolation. + - ``outlet_events`` is carried over by reference, not copied: where a sub-task's + events go is not decided here. ``IndexedTaskRunner.indexed_context`` replaces it + with an accessor of the sub-task's own, so that its events can be checkpointed + apart from its siblings', and ``IterableOperator`` merges that accessor into the + parent's once the sub-task has succeeded. + + Use cases + - Multithreading: when using thread-based executors (``concurrent.futures`` + ThreadPoolExecutor) multiple threads share memory; cloning prevents + concurrent mutation of shared structures. + - Async concurrency: when running coroutine-based tasks concurrently in + the same event loop, tasks may still mutate shared mappings; cloning + avoids interference. + - Multiprocessing: while processes do not share memory, cloning keeps the + semantics consistent and avoids accidentally capturing references that + would be pickled. + + Performance + - The implementation intentionally copies only a small set of commonly + mutated fields rather than performing a blanket deep copy of the entire + context to keep the operation cheap. If future code stores additional + mutable state in the context that needs isolation, this function should + be extended appropriately. + + :param context: The original execution context to clone. + :returns: A new :class:`Context` safe to hand to a concurrently running task. + + :meta private: + """ + cloned_context = Context() + cloned_context.update(context) + cloned_context["params"] = copy.deepcopy(context.get("params", {})) + cloned_context["inlets"] = list(context.get("inlets", [])) + cloned_context["outlets"] = list(context.get("outlets", [])) + templates_dict = cloned_context.get("templates_dict") + if templates_dict is not None: + cloned_context["templates_dict"] = copy.deepcopy(templates_dict) + # Everything else, dag_run and the event accessors included, came over by reference from + # update(); none of those keys has to be present. + return cloned_context + + def context_merge(context: Context, *args: Any, **kwargs: Any) -> None: """ Merge parameters into an existing context. diff --git a/task-sdk/src/airflow/sdk/definitions/iterableoperator.py b/task-sdk/src/airflow/sdk/definitions/iterableoperator.py new file mode 100644 index 0000000000000..57925c029fe51 --- /dev/null +++ b/task-sdk/src/airflow/sdk/definitions/iterableoperator.py @@ -0,0 +1,1292 @@ +# +# 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 + +import asyncio +import hashlib +import json +import os +import threading +import time +from collections.abc import AsyncIterable, AsyncIterator, Callable, Iterable, Mapping, Sequence +from functools import partial +from typing import TYPE_CHECKING, Any + +try: + # Python 3.11+ + BaseExceptionGroup +except NameError: + from exceptiongroup import BaseExceptionGroup + +from airflow.sdk import BaseXCom, TaskInstanceState, TriggerRule +from airflow.sdk.bases.operator import BaseOperator, event_loop +from airflow.sdk.bases.skipmixin import SkipMixin +from airflow.sdk.bases.xcom import XComIterable +from airflow.sdk.definitions.retry_policy import RetryAction, RetryDecision +from airflow.sdk.definitions.xcom_arg import XComArg +from airflow.sdk.exceptions import ( + AirflowFailException, + AirflowRescheduleException, + AirflowSensorTimeout, + AirflowSkipException, + AirflowTaskTerminated, + AirflowTaskTimeout, + DagRunTriggerException, + DownstreamTasksSkipped, + TaskDeferred, +) +from airflow.sdk.execution_time.comms import DeadlockImminentError +from airflow.sdk.execution_time.context import context_update_for_unmapped +from airflow.sdk.execution_time.executor import AsyncAwareExecutor +from airflow.sdk.execution_time.task_runner import ( + IndexedTaskInstance, + IndexedTaskRunner, + IndexedTaskState, + _push_xcom_if_needed, +) +from airflow.sdk.serde import serialize + +if TYPE_CHECKING: + import jinja2 + + from airflow.sdk.definitions._internal.expandinput import ExpandInput, Resolved + from airflow.sdk.definitions.context import Context + from airflow.sdk.definitions.mappedoperator import MappedOperator + from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance + from airflow.sdk.types import Logger + + +# The trigger rules under which one skipped upstream task instance skips a task, whatever the other +# upstream task instances did (see TriggerRuleDep). A skipped iteration has the same effect on a +# downstream task with one of these rules as a skipped mapped task instance would. +SKIPPED_WITH_A_SKIPPED_UPSTREAM = frozenset( + {TriggerRule.ALL_SUCCESS, TriggerRule.NONE_SKIPPED, TriggerRule.ALL_DONE_MIN_ONE_SUCCESS} +) + + +# Raised by an indexed task, these fail the task without a retry, as the runner does for a task that raises +# them itself (see _run_task_and_map_outcome and _handle_handler_failure). +FAIL_WITHOUT_RETRY = (AirflowFailException, AirflowSensorTimeout, AirflowTaskTerminated) + + +# How strongly a retry policy decision speaks for the task when several indexed tasks failed: one the +# policy says must not be retried fails the task, one it says to retry makes it retry on its terms. +class Checkpoints: + """ + Decide whether one attempt of an IterableOperator may resume from its per-index checkpoints. + + Checkpoints are only consulted from the second attempt onwards. A completion marker left by a + previous fully successful run means this attempt follows a manual clear (which raises + ``max_tries`` but does not reset ``try_number``), so every index must run again: the stale + ``SUCCESS`` checkpoints are ignored and overwritten. On entry the marker is replaced by the + attempt the rerun starts at, so that a crash during the rerun resumes from the checkpoints + written since, and only from those: an index the crashed rerun did not reach still holds its + checkpoint from before the clear, which must not be replayed. The marker is written again when + the block exits without an exception, and when every iteration skipped: the task is then + ``SKIPPED``, a final state, and a clear of it must run every iteration again rather than + replay the skips, as a cleared mapped task instance would. An attempt that fails, even with + some iterations skipped, writes no marker, so a retry or a clear after it resumes: iterations + that succeeded or skipped keep that outcome and only the others run again. + + The checkpoints themselves are never deleted: one marker write costs the same whatever the indexed task + count, and the store is scoped to the parent task instance, so state a sub-task stored for itself + is never touched. They expire with the store's default retention (``[state_store] + default_retention_days``, 30 days unless configured, 0 disables expiry), which also bounds how + long a task that exhausted its retries keeps them. Keeping them until then is intended: a manual + clear of such a task resumes from the checkpoints instead of re-running every index, and the + marker is what tells a clear-after-success apart from that. + """ + + # Same namespace as IndexedTaskState.build_key, for the same reason. + COMPLETION_KEY = "_iterable_completed" + + def __init__(self, context: Context) -> None: + self._store = context["task_state_store"] + self._try_number = context["ti"].try_number + self.trust_checkpoints = False + # The attempt from which checkpoints may be resumed; older ones predate a manual clear. + self.since = 0 + + def __enter__(self) -> Checkpoints: + if self._try_number > 1: + marker = self._store.get(self.COMPLETION_KEY) + if isinstance(marker, Mapping) and marker.get("completed"): + self._store.set(self.COMPLETION_KEY, {"completed": False, "since": self._try_number}) + else: + self.trust_checkpoints = True + since = marker.get("since") if isinstance(marker, Mapping) else None + if isinstance(since, int): + self.since = since + return self + + def __exit__(self, exc_type, exc_value, traceback) -> None: + if exc_type is None or issubclass(exc_type, AirflowSkipException): + self._store.set(self.COMPLETION_KEY, {"completed": True, "try_number": self._try_number}) + + +class IndexedTaskInstanceNotStarted(Exception): + """ + The outcome of an indexed task instance pulled before a kill and reached after it. + + Never raised: ``IterableOperator._run_task`` returns it as the instance's outcome, so that + :meth:`IndexedTaskOutcomes.record` counts it apart from the failures. It ran no code, wrote + no checkpoint and fires no callback; the next attempt runs it. + """ + + +class IterationState: + """ + The state of one run of an iterated task, kept apart from the operator's configuration. + + It holds what the run needs to remember while it is going: the sub-operators in flight and + those already killed, so that :meth:`IterableOperator.on_kill` reaches each one once; the + stop flag a kill sets, which the executor consults before starting the next indexed task; + and the resolved input. What the indexed tasks ended with is :class:`IndexedTaskOutcomes`. + + A copy of the operator (``copy.copy`` in ``prepare_for_execution``, the deep copy of + ``dag.partial_subset``) is another task with nothing in flight and no kill pending, so copying + the state gives a fresh one. The + operator renews it once a run ended, not when one starts: a rerun in the same process starts + clean, while a kill that arrives before the run sets the stop flag on the state the run uses. + """ + + def __init__(self) -> None: + # Keyed by identity: BaseOperator equality compares fields such as task_id, which every + # sub-operator of one iterated task shares, so a set would hold one of them at most. + self._in_flight: dict[int, BaseOperator] = {} + self._lock = threading.Lock() + # The runner calls on_kill() again after the execution timeout that made _run_tasks call + # it first, and a sub-operator is killed once. + self._killed: set[int] = set() + # A plain flag and a plain list, not an Event and not guarded by the lock: on_kill() runs in + # the runner's SIGTERM handler, on the main thread, between two bytecodes of whatever the + # loop thread was doing, register() or unregister() under the lock included. A lock the + # handler takes that its own thread holds would hang it, so request_stop() and start_kill() + # take none; take_in_flight() is called from the kill thread instead. + self._stop_requested = False + # The threads on_kill() started to kill the sub-operators in flight; see await_kill. + self._kill_threads: list[threading.Thread] = [] + #: The input resolved for this task instance, once ``aresolve`` returned. + self.resolved: Resolved | None = None + + def __copy__(self) -> IterationState: + return IterationState() + + def __deepcopy__(self, memo: dict[int, Any]) -> IterationState: + return IterationState() + + def register(self, operator: BaseOperator) -> None: + """Note that ``operator`` is executing, so a kill reaches it.""" + with self._lock: + self._in_flight[id(operator)] = operator + + def unregister(self, operator: BaseOperator) -> None: + """Note that ``operator`` is done, one way or another.""" + with self._lock: + self._in_flight.pop(id(operator), None) + + def __contains__(self, operator: object) -> bool: + with self._lock: + return id(operator) in self._in_flight + + def take_in_flight(self) -> list[BaseOperator]: + """Return the sub-operators in flight that were not handed out before, and mark them killed.""" + with self._lock: + operators = [op for key, op in self._in_flight.items() if key not in self._killed] + self._killed.update(map(id, operators)) + return operators + + def request_stop(self) -> None: + """Ask the iteration to start nothing else; see :meth:`stop_requested`. Takes no lock.""" + self._stop_requested = True + + def start_kill(self, kill: Callable[[], None]) -> None: + """ + Run ``kill`` in a thread of its own, kept so that :meth:`await_kill` can wait for it. + + Takes no lock, so it can be called from a signal handler; ``kill`` takes the sub-operators + in flight itself, with :meth:`take_in_flight`, on its thread. + """ + thread = threading.Thread(target=kill, name="iterable-operator-on-kill", daemon=True) + self._kill_threads.append(thread) + thread.start() + + def await_kill(self, timeout: float) -> None: + """ + Wait up to ``timeout`` seconds for the threads :meth:`start_kill` started. + + Called by ``IterableOperator._run_tasks`` once the loop closed and before the run + concludes, so the task does not end, and the process with it, while a sub-operator's + ``on_kill`` is still cleaning up. The threads are daemon threads: one that never returns + holds the run for ``timeout`` at most, and the supervisor's SIGKILL bounds the rest. + """ + deadline = time.monotonic() + timeout + for thread in list(self._kill_threads): + thread.join(max(0.0, deadline - time.monotonic())) + + def stop_requested(self) -> bool: + """Whether :meth:`request_stop` was called; passed to the executor as its ``stop``.""" + return self._stop_requested + + @property + def length(self) -> int | None: + """How many indexed tasks the resolved input gives, or None while it is still being resolved.""" + return self.resolved.length if self.resolved is not None else None + + +class IndexedTaskOutcomes: + """ + What the indexed tasks of one run ended with, and what the task ends with because of it. + + Entered by ``IterableOperator._run_tasks`` around the loop that drives the indexed tasks. + :meth:`record` takes each outcome: a skip is noted, an outcome the iteration cannot carry + raises at once, any other failure is logged and collected. :meth:`conclude` turns what was + collected, and whether the run's :class:`IterationState` was asked to stop, into the task's + own outcome and raises it: killed, failed on the exception the runner judges, skipped over + an empty input or because every indexed task skipped. Whatever + exception leaves the block, those included, the failed indexed tasks' callbacks are reported + on exit with the task's fate, so they say what the runner then does: retry, or fail for good. + """ + + #: How strongly a retry policy decision speaks for the task when several indexed tasks failed: + #: one the policy fails the task on outweighs one it retries on, which outweighs the default. + _DECISION_WEIGHT = {RetryAction.FAIL: 2, RetryAction.RETRY: 1, RetryAction.DEFAULT: 0} + + def __init__(self, operator: IterableOperator, state: IterationState, context: Context) -> None: + self._operator = operator + self._state = state + self._context = context + self.total = 0 + #: Pulled before a kill and never started (see :class:`IndexedTaskInstanceNotStarted`). + self.not_started = 0 + self.exceptions: list[BaseException] = [] + self.skipped: dict[int, AirflowSkipException] = {} + self._failed_runners: list[IndexedTaskRunner] = [] + self._decision: tuple[BaseException, RetryDecision] | None = None + + def __enter__(self) -> IndexedTaskOutcomes: + return self + + def __exit__(self, exc_type, exc_value, traceback) -> None: + # Whatever the task ends with decides every failed indexed task's callback, so they agree. + if exc_value is not None and self._failed_runners: + self._report(exc_value) + + def note_failed(self, runner: IndexedTaskRunner) -> None: + """Remember a failed indexed task's runner: its callback waits for the task's fate.""" + self._failed_runners.append(runner) + + @property + def log(self) -> Logger: + """The operator's logger: the outcomes are logged under the task they belong to.""" + return self._operator.log + + @property + def failed_runners(self) -> tuple[IndexedTaskRunner, ...]: + """The runners of the indexed tasks that failed in this run, in the order they failed.""" + return tuple(self._failed_runners) + + def record(self, task: IndexedTaskInstance, raised: BaseException | None) -> None: + """Take one indexed task's outcome: nothing, a skip, not started, a failure, or one ending it.""" + if isinstance(raised, IndexedTaskInstanceNotStarted): + self.not_started += 1 + return + self.total += 1 + if raised is None: + return + if isinstance(raised, AirflowSkipException): + self.skipped[task.index] = raised + return + # An outcome the iteration cannot carry stops it at once: every later indexed task would + # end the same way. + if (fail_fast := self._fail_fast_for(task, raised)) is not None: + raise fail_fast from raised + self.log.exception( + "An exception occurred for task_id %s with index %s", task.task_id, task.index, exc_info=raised + ) + self.exceptions.append(raised) + + def conclude(self) -> list[int]: + """ + Raise the task's outcome from what was recorded, or return the skipped indices. + + A kill comes first: with every killed indexed task returning normally there would be no + failure to raise, the completion marker would be written and a result returned over + XComs that do not exist, and a kill while the input resolves is not an empty input. The + task then fails without a retry, as the runner treats a terminated task, and no marker + is written, so a later clear resumes from the checkpoints of the indexed tasks that did + finish. Then the failures, handed to the runner as the exception it judges; then an empty + input, skipped as a mapped task over nothing is; then every indexed task skipped, which + skips the task too. + """ + if self._state.stop_requested(): + raise AirflowTaskTerminated(self._killed_message()) from ( + BaseExceptionGroup("Sub-task failures", self.exceptions) if self.exceptions else None + ) + if self.exceptions: + raise self._failure_for_the_runner() + if self.total == 0: + raise AirflowSkipException("The input to iterate over is empty.") + if self.skipped and len(self.skipped) == self.total: + raise next(iter(self.skipped.values())) + return sorted(self.skipped) + + def _killed_message(self) -> str: + """ + Say what a kill left behind, naming only the counts that are not zero. + + Before the input was resolved there is no count to give. Once it is, the indexed tasks + that ran (the killed ones among them, which came back as failures), those pulled before + the kill and never started, and those never pulled add up to the length. + """ + length = self._state.length + if length is None: + return "The iterated task was killed before its input was resolved." + parts = [f"{self.total} of {length} items ran"] + if self.not_started: + parts.append(f"{self.not_started} pulled but never started") + if never_pulled := length - self.total - self.not_started: + parts.append(f"{never_pulled} never pulled") + return f"The iterated task was killed: {', '.join(parts)}." + + def _report(self, raised: BaseException) -> None: + """ + Report each failed indexed task as what happens to the task: retried, or failed for good. + + Their callbacks were held back until every indexed task had run, so the retry callback of + one no longer announces a retry a sibling's ``AirflowFailException`` then rules out. They + run one after another, on the thread that ran the iteration. + """ + task_will_retry = self._task_will_retry(raised) + for runner in self._failed_runners: + runner.report_failure(task_will_retry=task_will_retry) + + def _task_will_retry(self, raised: BaseException) -> bool: + """ + Whether the runner retries the task for ``raised``, by the runner's own rules. + + No retry for the fail-fast exceptions nor for what is not an ``Exception`` (other than the + parent's timeout), none when the retry policy decides FAIL, and otherwise a retry while the + parent has attempts left (``IndexedTaskInstance.is_eligible_to_retry``). The policy's + decision is the one :meth:`_failure_for_the_runner` took for ``raised`` when it chose it, + so the policy is not evaluated again for the callbacks: a policy that calls a model may + answer differently each time, and the callbacks must say what the exception handed over + says. + """ + if isinstance(raised, FAIL_WITHOUT_RETRY) or not isinstance(raised, (Exception, AirflowTaskTimeout)): + return False + if not self._failed_runners[0].task_instance.is_eligible_to_retry: + return False + if (policy := self._operator.retry_policy) is not None: + if self._decision is not None and self._decision[0] is raised: + decision = self._decision[1] + else: + ti = self._context["ti"] + from_server = getattr(ti, "_ti_context_from_server", None) + max_tries = from_server.max_tries if from_server else ti.max_tries + try: + decision = policy.evaluate( + exception=raised, try_number=ti.try_number, max_tries=max_tries, context=self._context + ) + except Exception: + return True + if decision.action == RetryAction.FAIL: + return False + return True + + @staticmethod + def _fail_fast_for(task: IndexedTaskInstance, raised: BaseException) -> AirflowFailException | None: + """ + Return the failure that ends the task at once for an outcome the iteration cannot carry. + + An indexed task is not a task instance of its own: it cannot defer, reschedule, trigger a + DAG run or skip downstream tasks, and a ``BaseException`` that is no ``Exception`` + (``DeadlockImminentError``, ``KeyboardInterrupt``, ``SystemExit``) must never be swallowed + into a retry. Each of these fails the whole task without a retry, with a message saying + why. Any other exception, an ``AirflowTaskTimeout`` the operator raised itself included, is + the indexed task's own failure, which :meth:`record` collects. + """ + sub_task = f"Sub-task {task.task_id}[{task.index}]" + if isinstance(raised, TaskDeferred): + return AirflowFailException( + f"{sub_task} attempted to defer. Deferrable operators are not supported inside IterableOperator." + ) + if isinstance(raised, (DagRunTriggerException, DownstreamTasksSkipped)): + return AirflowFailException( + f"{sub_task} raised {type(raised).__name__}. Triggering DAG runs (TriggerDagRunOperator) and " + "skipping downstream tasks (ShortCircuitOperator and similar) are not supported inside " + "IterableOperator: the sub-task's index has no downstream tasks or DAG run of its own for " + "the effect to apply to." + ) + if isinstance(raised, AirflowRescheduleException): + return AirflowFailException( + f"{sub_task} attempted to reschedule (raised AirflowRescheduleException). Reschedule-mode " + "sensors are not supported inside IterableOperator: the sub-task's index has no task " + "instance of its own to reschedule." + ) + if isinstance(raised, DeadlockImminentError): + return AirflowFailException( + f"{sub_task} made a synchronous SDK call (e.g. Variable.get, BaseHook.get_connection/get_hook, " + "ti.xcom_pull) on the event loop thread while another sub-task's async SDK call was in " + "flight, which would deadlock the loop, so it is detected and raised eagerly instead. " + "Inside IterableOperator only async sub-tasks run on that thread: an async operator's " + "aexecute(), its pre_execute/post_execute and its callbacks. Use the async-safe " + "equivalents there (e.g. Variable.aget/aset, Hook.aget_connection/aget_hook, " + "ti.axcom_pull); sync sub-tasks and their callbacks run in worker threads, where the " + "same calls wait their turn." + ) + if isinstance(raised, AirflowTaskTimeout): + # The parent's timeout never gets here (_run_task re-raises it); this one the operator + # raised itself, off the main thread, and the runner judges it as it does for a plain + # task that raises it: a failure it retries. + return None + if not isinstance(raised, Exception): + return AirflowFailException( + f"{sub_task} raised a non-Exception BaseException: {type(raised).__name__}: {raised}" + ) + return None + + def _failure_for_the_runner(self) -> BaseException: + """ + Pick the exception the runner decides the task's outcome on. + + The runner classifies by exception type: a fail-fast exception fails without a retry, and a + ``retry_policy`` matches rules against the type. A ``BaseExceptionGroup`` defeats both, so + an indexed task's own exception is handed over whenever one decides: the first fail-fast + one, the only one, or the one whose policy decision weighs most. With several failures the + others stay attached as its cause, so every traceback reaches the log. Several failures no + policy decides between are raised as a group, which the task's own retries then apply to. + The policy is evaluated once per failure, and the decision for the chosen one is kept for + the callbacks, as is the default for a group no decision outweighs (see + :meth:`_task_will_retry`). + """ + exceptions = self.exceptions + group = BaseExceptionGroup("Multiple sub-task failures", exceptions) + chosen: BaseException | None = next( + (exc for exc in exceptions if isinstance(exc, FAIL_WITHOUT_RETRY)), None + ) + if chosen is None and len(exceptions) == 1: + return exceptions[0] + if chosen is None and (policy := self._operator.retry_policy) is not None: + ti = self._context["ti"] + from_server = getattr(ti, "_ti_context_from_server", None) + max_tries = from_server.max_tries if from_server else ti.max_tries + weights = [] + decisions: list[RetryDecision | None] = [] + for exc in exceptions: + try: + decision = policy.evaluate( + exception=exc, try_number=ti.try_number, max_tries=max_tries, context=self._context + ) + weights.append(self._DECISION_WEIGHT.get(decision.action, 0)) + decisions.append(decision) + except Exception: + # As the runner does: a policy that fails to evaluate leaves the default. + self.log.exception("Retry policy evaluation failed for a sub-task failure") + weights.append(0) + decisions.append(None) + if max(weights) > 0: + index = weights.index(max(weights)) + chosen = exceptions[index] + if (kept := decisions[index]) is not None: + self._decision = (chosen, kept) + else: + # Every decision was the default, or failed to evaluate, which the runner treats + # the same: the group carries the default, so the callbacks follow it instead of + # evaluating the policy once more, on an exception no item raised. + self._decision = (group, RetryDecision.default()) + if chosen is None: + return group + if len(exceptions) > 1: + chosen.__cause__ = group + return chosen + + +class IterableOperator(BaseOperator): + """ + Operator used for Iterable Tasks (IT) that runs a mapped operator over an iterable input. + + The IterableOperator wraps a :class:`MappedOperator` together with an + :class:`ExpandInput` and is responsible for creating and running the + per-index runtime task instances. The IterableOperator itself participates + in Airflow's native retry mechanism — its ``retries`` and ``retry_delay`` + are inherited from the wrapped operator so that when any sub-task needs + a retry the whole IterableOperator is retried by Airflow. Already-succeeded + sub-tasks are skipped on each retry attempt because their state is + checkpointed in the ``task_state_store``. + + The IterableOperator executes the mapped operator instances using a + concurrent executor with a configurable number of workers. By default + the worker count is taken from the mapped operator's ``partial_kwargs`` + (``task_concurrency``) if present, otherwise falls back to + ``os.cpu_count()`` and finally to ``1``. ``os.cpu_count()`` counts the CPUs of + the machine: a worker in a container with a CPU limit still sees every CPU of + its node, so set ``task_concurrency`` there. + + **Crash recovery:** When the worker crashes mid-iteration and the task is re-run (e.g. via a + manual clear), already-succeeded sub-tasks are skipped and only the pending/failed ones are + executed again. Every sub-task inherits its ``try_number`` from the IterableOperator's own task + instance, so the attempt count reported to a sub-task matches the attempt Airflow is currently + running. The checkpoint is only consulted from the second attempt onwards, and solely to decide + whether an index already succeeded. Once every index has succeeded, a completion marker is written + so that a *subsequent* manual clear (which does not reset ``try_number``) re-runs every index from + scratch instead of replaying the previous run's stale results (see :class:`Checkpoints`). A + checkpoint carries the indexed task's full result (plus its extra XComs and outlet events) and lives in + the ``task_state_store`` table for the store's retention, unless a ``[workers] + state_store_backend`` is configured and only a reference is stored; a custom XCom backend alone + does not keep large results out of the metadata database. + + :param operator: The :class:`MappedOperator` to unmap and execute for + each element of ``expand_input``. Each indexed runtime receives a + deep copy/unmapped instance of this operator. + + :param expand_input: Provider of the values to iterate + over. Its ``aresolve(context)`` method gives the indexed task count and the + per-index ``mapped_kwargs`` used to unmap the operator. + + :param kwargs: Additional keyword arguments forwarded to + :class:`BaseOperator` when instantiating the IterableOperator + (e.g. ``dag``, ``start_date``). + + :returns: An :class:`XComIterable` if the mapped operator pushes XComs, otherwise ``None``. + + .. note:: + ``multiple_outputs`` is ignored for iterated tasks. Each sub-task's return value is pushed + whole as ``return_value_`` and the task's own return value is the ``XComIterable`` + over them, so a ``Mapping`` return annotation on the wrapped ``@task`` does not fan its + keys out into separate XComs the way it does for ``.expand()``. + + .. note:: + Deferred operators (those that raise :class:`~airflow.sdk.exceptions.TaskDeferred`) are not + supported yet inside IterableOperator. A ``TaskDeferred`` exception raised by an indexed task + instance will propagate as an error rather than pausing and resuming the task. + + Reschedule-mode sensors (those that raise :class:`~airflow.sdk.exceptions.AirflowRescheduleException`) + are also not supported. A reschedule raised by an indexed task instance will fail the whole + IterableOperator immediately with a clear error rather than being silently mishandled. + + Triggering DAG runs (:class:`~airflow.sdk.exceptions.DagRunTriggerException`, raised by + ``TriggerDagRunOperator``) and skipping downstream tasks are not supported either: a sub-task + index has no DAG run or downstream tasks of its own for the trigger/skip to apply to. + Operators that can skip downstream tasks (``ShortCircuitOperator``, the branch operators, + ``@task.short_circuit``, ``@task.branch`` and any other ``SkipMixin``) are rejected by + ``.iterate()`` itself, since inside an iteration they would find nothing to skip and let + every downstream task run. A trigger, or a + :class:`~airflow.sdk.exceptions.DownstreamTasksSkipped` raised anyway, fails the whole + IterableOperator immediately with a clear error rather than silently doing nothing. + + Sub-task outcomes are classified before being aggregated: if any sub-task raises + :class:`~airflow.sdk.exceptions.AirflowFailException`, that exception is re-raised directly so + the IterableOperator fails without retrying. A sub-task that raises + :class:`~airflow.sdk.exceptions.AirflowSkipException` is skipped, as a mapped task instance + would be: it pushes no XCom, does not fail the task and is not run again on a retry. It is + left out of the task's :class:`~airflow.sdk.bases.xcom.XComIterable`, so downstream tasks + only see the values that exist. A sub-task that returns ``None`` is not skipped: it succeeds, + pushes no XCom and keeps its position in the sequence, which reads ``None`` there, so a + downstream ``.expand()`` over the result runs over it; ``.expand()`` leaves a mapped task + instance that returned ``None`` out of the count instead. A direct downstream task whose trigger rule skips it when an + upstream task instance is skipped (``all_success``, ``none_skipped``, + ``all_done_min_one_success``) is skipped, as after a mapped upstream; one with a rule such + as ``none_failed`` runs over the remaining values. If *every* sub-task is skipped, a single + ``AirflowSkipException`` is re-raised so the IterableOperator itself is marked ``SKIPPED``, and + so is it over an empty input, as a mapped task over nothing is. All other sub-task exceptions are aggregated + into a :class:`BaseExceptionGroup` and treated as a regular retryable failure. + + .. warning:: + **Inputs are shared between iterations.** + + All iterations run in one process, so a value handed to several of them is the same object + in each of them: every value passed through ``.partial()``, and with + ``.iterate(a=..., b=...)`` every element of ``a`` and of ``b``, which the cross product + combines more than once. A mapped task instance gets its own copy, because it runs in its + own process; an iteration does not. Treat inputs as read-only, or copy what the task + changes in place. + + .. note:: + **Callbacks run per indexed task, and a failed one's wait for the task's fate.** + + ``on_success_callback`` and ``on_skipped_callback`` run once the indexed task's checkpoint is + written, so they speak for work a retry will not run again, where it ran: in its worker + thread for a sync operator, on the event loop for an async one. A failed indexed task's + ``on_failure_callback`` or ``on_retry_callback`` runs once every indexed task has run, on the + thread that ran the iteration, and says what happens to the task: retried or failed for good + (see :class:`IndexedTaskOutcomes`). A failure no indexed task + owns, such as an error resolving the input, fires no callback, and neither does an indexed + task pulled before a kill and never started: the iterated task has no callbacks of its own. + Listeners fire once, for the task instance, as for any task; an indexed task is not a task + instance and fires none. + + .. note:: + **Pools count the task instance, not its iterations.** + + The scheduler reserves ``pool_slots`` once for the iterated task, while up to + ``task_concurrency`` iterations run inside it. A pool sized to cap the load on a shared + resource (database connections, the rate limit of an API) therefore sees one reservation + for that many concurrent uses. Choose ``task_concurrency`` with the pool in mind; reserving + slots per iteration needs support in the scheduler, which does not exist yet. + + .. warning:: + **Async sub-tasks must only make async SDK calls.** + + IterableOperator runs multiple async sub-tasks concurrently on the same event loop, each + making async SDK calls of its own (checkpointing, XCom push). If an async sub-task's + ``aexecute()`` — or a hook/callback it calls — issues a *synchronous* SDK call instead (e.g. + ``Variable.get``, ``BaseHook.get_connection``/``get_hook``, ``ti.xcom_pull``, or a sync + ``on_success_callback``/``pre_execute``), it can collide with another sub-task's async SDK + call that is concurrently holding the communication lock, which is detected and raised + eagerly as a non-retryable failure rather than silently deadlocking. Use the async-safe + equivalents inside async operators: :meth:`~airflow.sdk.bases.hook.BaseHook.aget_connection`/ + ``aget_hook``, ``ti.axcom_pull`` and ``Variable.aget``/``aset``. Sync sub-tasks are not + concerned: they run in worker threads, their ``execute``, hooks and callbacks included, + where a synchronous SDK call waits for the lock; so does ``on_kill`` of the sub-operators, + which :meth:`on_kill` runs off the loop thread. + + .. warning:: + **``execution_timeout`` caps the whole iteration, not each indexed task.** + + The IterableOperator keeps the wrapped operator's ``execution_timeout`` as a wall-clock limit + on the entire task instance. The runner enforces it on the main thread exactly as for any + other task, so an iteration that overruns fails with ``AirflowTaskTimeout`` and + :meth:`on_kill` is propagated to every sub-task still in flight. Since ``.iterate()`` runs + all indexed tasks in one task instance, this is the per-instance limit of ``.expand()`` applied + to the whole iteration rather than to each indexed task: no limit applies per indexed task, + sync or async (one of the same value, started later, could never fire first). An + ``AirflowTaskTimeout`` a sync sub-task raises itself (a hook that gave up waiting) is that + sub-task's own failure, as it is the mapped task instance's under ``.expand()``. + """ + + _operator: MappedOperator + expand_input: ExpandInput + partial_kwargs: dict[str, Any] + shallow_copy_attrs: Sequence[str] = ( + "_operator", + "expand_input", + "partial_kwargs", + "_log", + ) + + @staticmethod + def _refuse_operators_that_skip_downstream(operator: MappedOperator) -> None: + """ + Refuse to iterate an operator that can skip downstream tasks. + + An iteration has no downstream tasks of its own, so ``ShortCircuitOperator``, the branch + operators and any other ``SkipMixin`` would skip nothing and let every downstream task run. + Checked on the class: ``MappedOperator._can_skip_downstream`` is only derived from ``SkipMixin`` + on the classic path, while the ``@task`` path copies a class default that is ``False`` even for + ``@task.short_circuit`` and ``@task.branch``. + """ + if issubclass(operator.operator_class, SkipMixin): + raise TypeError( + f"{operator.operator_name} can skip downstream tasks and cannot be iterated: an iteration " + f"of {operator.task_id!r} has no downstream tasks of its own, so it would skip nothing and " + "every downstream task would run. Use .expand() for it instead." + ) + + @staticmethod + def _unprefixed_task_id(operator: MappedOperator) -> str: + """ + Return the wrapped operator's task id without its task group's prefix. + + ``partial()`` already gave the wrapped operator the prefixed id, and ``BaseOperator.__init__`` + prefixes the id it gets once more, since an IterableOperator is not built from a mapped + operator. Handing it the bare id keeps the two equal. The same rule as ``label``, which cannot + be used here because it returns the display name when there is one. + """ + task_group = operator.task_group + if task_group and task_group.node_id and task_group.prefix_group_id: + return operator.task_id[len(task_group.node_id) + 1 :] + return operator.task_id + + def __init__( + self, + *, + operator: MappedOperator, + expand_input: ExpandInput, + **kwargs, + ): + if operator.get_closest_mapped_task_group() is not None: + raise NotImplementedError("operator expansion in an expanded task group is not yet supported") + self._refuse_operators_that_skip_downstream(operator) + + super().__init__( + **{ + **kwargs, + "task_id": self._unprefixed_task_id(operator), + "owner": operator.owner, + "email": operator.email, + "email_on_retry": operator.email_on_retry, + "email_on_failure": operator.email_on_failure, + "retries": operator.retries, + "retry_delay": operator.retry_delay, + "retry_exponential_backoff": operator.retry_exponential_backoff, + "max_retry_delay": operator.max_retry_delay, + "retry_policy": operator.retry_policy, + "start_date": operator.start_date, + "end_date": operator.end_date, + "depends_on_past": operator.depends_on_past, + "ignore_first_depends_on_past": operator.ignore_first_depends_on_past, + "wait_for_past_depends_before_skipping": operator.wait_for_past_depends_before_skipping, + "wait_for_downstream": operator.wait_for_downstream, + "dag": operator.dag, + "params": operator.params, + "priority_weight": operator.priority_weight, + "weight_rule": operator.weight_rule, + "queue": operator.queue, + "pool": operator.pool, + "pool_slots": operator.pool_slots, + # Kept as the wall-clock cap on the whole iteration, enforced by the runner (see the + # class docstring); no limit applies per indexed task. + "execution_timeout": operator.execution_timeout, + "trigger_rule": operator.trigger_rule, + "resources": operator.resources, + "run_as_user": operator.run_as_user, + "map_index_template": operator.map_index_template, + "max_active_tis_per_dag": operator.max_active_tis_per_dag, + "max_active_tis_per_dagrun": operator.max_active_tis_per_dagrun, + "executor": operator.executor, + "executor_config": operator.executor_config, + "do_xcom_push": operator.partial_kwargs.get("do_xcom_push", True), + # Ignored for iterated tasks (also when passed explicitly): the return value pushed by the + # runner is the XComIterable aggregate, not a dict, and every sub-task result is pushed + # whole under return_value_. The wrapped @task may still infer True from a + # Mapping return annotation, which would make the runner reject the aggregate. + "multiple_outputs": False, + "inlets": operator.inlets, + "outlets": operator.outlets, + "task_group": operator.task_group, + "doc": operator.doc, + "doc_md": operator.doc_md, + "doc_json": operator.doc_json, + "doc_yaml": operator.doc_yaml, + "doc_rst": operator.doc_rst, + "task_display_name": operator.task_display_name, + "allow_nested_operators": operator.allow_nested_operators, + # The iterated task has no callbacks and no execute hooks of its own: they run per + # indexed task, from the wrapped operator's partial kwargs. Passed explicitly, since + # _apply_defaults would otherwise fill them from the DAG's default_args and the + # runner would run them for the task on top of the indexed tasks' own. + "on_execute_callback": None, + "on_success_callback": None, + "on_failure_callback": None, + "on_retry_callback": None, + "on_skipped_callback": None, + "pre_execute": None, + "post_execute": None, + } + ) + self._operator = operator + self.expand_input = expand_input + self.partial_kwargs = dict(operator.partial_kwargs) if operator.partial_kwargs else {} + task_concurrency = self.partial_kwargs.pop("task_concurrency", None) + if task_concurrency is not None and task_concurrency < 1: + raise ValueError(f"task_concurrency must be at least 1, got {task_concurrency}") + # pool_slots is reserved once for the task instance, not per iteration: see the class docstring. + self.max_workers = task_concurrency if task_concurrency is not None else (os.cpu_count() or 1) + # unmap() would normally apply these three flags to each generated sub-operator, and + # __attrs_post_init__ would apply them (plus the upstream-relationship wiring below) to the + # MappedOperator itself; since IterableOperator skips __attrs_post_init__ entirely (it isn't a + # MappedOperator), it must reproduce that part of the contract for its own single DAG node. + self.is_setup = bool(self.partial_kwargs.get("is_setup", False)) + self.is_teardown = bool(self.partial_kwargs.get("is_teardown", False)) + on_failure_fail_dagrun = self.partial_kwargs.get("on_failure_fail_dagrun", False) + if on_failure_fail_dagrun: + self.on_failure_fail_dagrun = on_failure_fail_dagrun + XComArg.apply_upstream_relationship(self, self.expand_input.value) + # Mirrors MappedOperator.__attrs_post_init__: partial kwargs corresponding to the wrapped + # operator's own template fields may themselves be XComArgs (e.g. `.partial(some_field=xcom)`), + # and those upstream edges must be recorded too, not just the ones from expand_input. + for key, value in self.partial_kwargs.items(): + if key in self._operator.template_fields: + XComArg.apply_upstream_relationship(self, value) + # What one run remembers while it is going (see IterationState); fresh for every run and + # for every copy of the operator. + self._state = IterationState() + + def __copy__(self) -> IterableOperator: + # prepare_for_execution copies the operator with copy.copy, which would share the state of + # the run with the Dag's operator: a kill in one attempt would then stop the next one under + # dag.test(). The copy is another run and starts with a state of its own (see IterationState). + other = type(self).__new__(type(self)) + other.__setstate__({**self.__getstate__(), "_state": IterationState()}) + return other + + # How long the run waits for its threads to end once it is over: the executor's worker and + # coroutine shutdown, and the thread on_kill() kills the sub-operators from. + _SHUTDOWN_TIMEOUT: float = 10.0 + + def on_kill(self) -> None: + # The default BaseOperator.on_kill() is a no-op, which would otherwise leave every + # currently in-flight sub-task unaware that the IterableOperator itself was killed + # (SIGTERM) or hit its execution_timeout: propagate to each active sub-operator instead. + # First stop the iteration from starting anything else: the killed indexed tasks come back + # as failures and free their slots, which would otherwise be filled with the next ones. + state = self._state + state.request_stop() + # Always in a thread of its own, which takes the sub-operators in flight itself. The + # runner's SIGTERM handler calls this on the main thread, between two bytecodes of the loop + # thread: a lock taken here that register() or unregister() holds at that moment would + # hang the handler, so nothing here takes one (see IterationState). On that thread the + # event loop either runs, and a synchronous SDK call in a sub-operator's on_kill would + # raise DeadlockImminentError, or is paused between two run_until_complete calls while a + # result is handed to the consumer, and the same call would wait for a lock a parked asend + # holds, which only the paused loop can release. In its own thread the call waits its turn + # in both cases, and the loop goes on serving the sub-tasks. _run_tasks kills what is in + # flight through _kill directly, from a thread the loop drives, and waits for this thread + # before the run concludes (see IterationState.await_kill). + state.start_kill(lambda: self._kill(state.take_in_flight())) + + def _kill(self, operators: list[BaseOperator]) -> None: + # One sub-operator's on_kill must not keep the kill from the others: DeadlockImminentError + # is a BaseException, so it is caught here as a plain error is. + for operator in operators: + try: + operator.on_kill() + except BaseException: + self.log.exception("Error calling on_kill() for sub-task operator %s", operator.task_id) + + @property + def returns_dag_result(self) -> bool: + return self._operator.returns_dag_result + + @returns_dag_result.setter + def returns_dag_result(self, value: bool) -> None: + self._operator.returns_dag_result = value + + @property + def operator_name(self) -> str: + # Shown as the wrapped operator (its class name, or a @task callable's custom_operator_name). + # task_type is not forwarded: it names the class that runs, and what resolves a class from it + # (the task's class reference, OpenLineage's extractors) would otherwise get the wrapped + # operator's, whose attributes this one does not have. + return self._operator.operator_name + + @property + def task_retries(self) -> int: + return self._operator.retries or 0 + + def _do_render_template_fields( + self, + parent: Any, + template_fields: Iterable[str], + context: Context, + jinja_env: jinja2.Environment, + seen_oids: set[int], + ) -> None: + # IterableOperator doesn't need to render template fields as the actual operator's template fields + # will be rendered in the IndexedTaskRunner when running each mapped task instance. + pass + + def _get_specified_expand_input(self) -> ExpandInput: + return self.expand_input + + def _render_unmapped_operator( + self, context: Context, unmapped_task: BaseOperator, jinja_env: jinja2.Environment + ) -> None: + context_update_for_unmapped(context, unmapped_task) + + unmapped_task._do_render_template_fields( + parent=unmapped_task, + template_fields=self._operator.template_fields, + context=context, + jinja_env=jinja_env, + seen_oids=set(), + ) + + async def axcom_push(self, task: IndexedTaskInstance, value: Any) -> None: + await task.axcom_push(key=BaseXCom.XCOM_RETURN_KEY, value=value) + + def _run_tasks( + self, + context: Context, + tasks: AsyncIterable[IndexedTaskInstance], + ) -> tuple[bool, list[int]]: + """Run ``tasks`` to completion; return whether they pushed results, and the skipped indices.""" + do_xcom_push = True + + self.log.info("Running tasks with %d workers", self.max_workers) + with ( + Checkpoints(context) as checkpoints, + IndexedTaskOutcomes(self, self._state, context) as outcomes, + ): + try: + with ( + event_loop() as loop, + AsyncAwareExecutor( + loop=loop, max_workers=self.max_workers, shutdown_timeout=self._SHUTDOWN_TIMEOUT + ) as executor, + ): + try: + for task, _result, raised in executor.imap_unordered( + partial( + self._run_task, + executor, + context, + trust_checkpoints=checkpoints.trust_checkpoints, + since=checkpoints.since, + outcomes=outcomes, + ), + tasks, + stop=self._state.stop_requested, + ): + do_xcom_push = task.do_xcom_push + outcomes.record(task, raised) + except BaseException: + # Whatever ends the loop early (the parent's execution_timeout, a failure that + # stops the task) is followed by the executor cancelling the coroutines, which + # would leave nothing registered for on_kill(); kill what is in flight first, + # off the loop thread, so that a sub-operator's synchronous SDK call in on_kill + # waits for the sub-tasks' calls in flight instead of raising. Awaited here, + # unlike on_kill()'s own thread: the loop keeps running meanwhile. + self._state.request_stop() + in_flight = self._state.take_in_flight() + try: + loop.run_until_complete(asyncio.to_thread(self._kill, in_flight)) + except RuntimeError: + self._kill(in_flight) + raise + finally: + # The kill on_kill() started runs on its own thread and may still be cleaning a + # sub-operator up once the killed indexed tasks came back and the loop closed. Waited + # for here, before the failure callbacks fire and the run concludes, so the task does + # not end with that cleanup half done. The loop is closed, so a synchronous SDK call + # in that on_kill no longer has an asend to wait for. + self._state.await_kill(self._SHUTDOWN_TIMEOUT) + skipped = outcomes.conclude() + return do_xcom_push, skipped + + def _skip_downstream_of_a_partial_skip(self, context: Context, result: XComIterable | None) -> None: + """ + Skip the downstream tasks a skipped mapped task instance would have skipped. + + Some iterations skipped and the others succeeded, so the task succeeds. Downstream tasks whose + trigger rule is satisfied by that are left to run over the values that exist; those whose + rule skips them when any upstream task instance skipped are skipped here, as they would be + after a mapped upstream. ``DownstreamTasksSkipped`` ends ``execute`` before the runner + pushes the return value, so it is pushed here first, for the downstream tasks that do run. + """ + to_skip = [ + task.task_id + for task in self.downstream_list + if task.trigger_rule in SKIPPED_WITH_A_SKIPPED_UPSTREAM + ] + if not to_skip: + return + ti = context["ti"] + if TYPE_CHECKING: + assert isinstance(ti, RuntimeTaskInstance) + _push_xcom_if_needed(result, ti, self.log) + raise DownstreamTasksSkipped(tasks=to_skip) + + async def _run_task( + self, + executor: AsyncAwareExecutor, + context: Context, + task: IndexedTaskInstance, + trust_checkpoints: bool, + since: int = 0, + *, + outcomes: IndexedTaskOutcomes, + ) -> tuple[IndexedTaskInstance, Any | None, BaseException | None]: + # See _run_tasks: checkpoints are consulted only on a retry that is not a rerun after a clear. + indexed_task_state = await task.aget_state() if trust_checkpoints else None + if indexed_task_state is not None and ( + indexed_task_state.try_number < since or indexed_task_state.fingerprint != task.input_fingerprint + ): + self.log.info( + "Running task instance %s for %s again: its checkpoint is from before the task was " + "cleared or for another input", + task.index, + task.task_id, + ) + indexed_task_state = None + if indexed_task_state is not None and indexed_task_state.status == TaskInstanceState.SUCCESS: + self.log.info( + "Skipping task instance %s for %s which already succeeded on a previous attempt", + task.index, + task.task_id, + ) + if indexed_task_state.result is not None: + await self.axcom_push(task, indexed_task_state.result) + for key, value in (indexed_task_state.xcoms or {}).items(): + await task.axcom_push(key=key, value=value) + if indexed_task_state.outlet_events: + indexed_task_state.replay_outlet_events(context["outlet_events"]) + return task, None, None + if indexed_task_state is not None and indexed_task_state.status == TaskInstanceState.SKIPPED: + return ( + task, + None, + AirflowSkipException( + f"Sub-task {task.task_id}[{task.index}] was skipped on a previous attempt" + ), + ) + + # Pulled before the kill and reached after it. on_kill() reaches the operators that started + # and the stop flag keeps the executor from pulling more; an item is neither until its runner + # starts, so a kill that lands while its checkpoint is read above would otherwise start it + # after the kill. A sync item is still handed to the pool below, whose threads are as many + # as the calls in flight, so only that pickup remains. + if self._state.stop_requested(): + return ( + task, + None, + IndexedTaskInstanceNotStarted( + f"Sub-task {task.task_id}[{task.index}] was pulled before the kill and never started" + ), + ) + + # The sub-task runs against its own view of the context (see IndexedTaskRunner.indexed_context), + # with its own outlet events: sub-tasks run concurrently and each needs its events + # attributed correctly so they can be checkpointed and merged individually (see + # IndexedTaskState.record_outlet_events and IndexedTaskRunner.merge_outlet_events_into). + indexed_task_runner = IndexedTaskRunner( + task_instance=task, + register=self._state, + ) + try: + if task.is_async: + with indexed_task_runner: + result = await indexed_task_runner.arun(context) + else: + # Entered and exited in the worker thread with execute, so what the exit does runs + # there too, as the reports below do (see _report_item). + + def run_indexed_task(): + with indexed_task_runner: + return indexed_task_runner.run(context) + + result = await executor.run_sync(run_indexed_task) + + indexed_task_state = IndexedTaskState( + status=TaskInstanceState.SUCCESS, + fingerprint=task.input_fingerprint, + try_number=task.try_number, + ) + # The result is checkpointed as well as pushed to XCom: the runner deletes every XCom + # key listed by the server before each attempt (xcom_keys_to_clear), so XCom alone + # cannot survive a retry, while the state store does. Both writes happen here, inside + # the coroutine, so they overlap with the sub-tasks still running instead of adding a + # synchronous pass over every index once the executor has drained. With a + # state_store_backend configured the checkpoint holds only a reference to the payload. + if result is not None and task.do_xcom_push: + indexed_task_state.result = result + indexed_task_state.record_outlet_events(indexed_task_runner.outlet_events) + # Written with the one checkpoint, not per push: a retry that skips this sub-task pushes + # them again, as the runner has deleted them by then. + if task.pushed_xcoms: + indexed_task_state.xcoms = dict(task.pushed_xcoms) + await task.aset_state(indexed_task_state) + except (asyncio.CancelledError, AirflowTaskTimeout) as stopped: + if isinstance(stopped, AirflowTaskTimeout) and indexed_task_runner.timed_out_on_its_own: + # Raised by the operator in its worker thread, which the parent's signal never + # reaches: this sub-task's own failure, as it is the mapped task instance's under + # .expand(); the siblings go on. + return await self._record_failure(task, indexed_task_runner, outcomes, stopped) + # Not this sub-task's outcome: it is being stopped from outside, by the executor + # cancelling it or by the parent's execution_timeout, whose signal handler raises on the + # main thread in whichever sub-task happens to run there. Both go on unchanged, so a + # cancellation stays one and the timeout reaches the runner, which retries the task. The + # sub-task the timeout struck is reported with the others (a cancelled one has nothing). + if isinstance(stopped, asyncio.CancelledError): + # A sync sub-task's thread may still be running: it gets no checkpoint from here on, + # so its exit must report nothing either (see IndexedTaskRunner.cancel). + indexed_task_runner.cancel() + if indexed_task_runner.failure is not None: + outcomes.note_failed(indexed_task_runner) + raise + except AirflowSkipException as e: + await task.aset_state( + IndexedTaskState( + status=TaskInstanceState.SKIPPED, + fingerprint=task.input_fingerprint, + try_number=task.try_number, + ) + ) + await self._report_item(executor, task, indexed_task_runner.report_skip) + return task, None, e + except BaseException as e: + return await self._record_failure(task, indexed_task_runner, outcomes, e) + + # The work is done and checkpointed: from here on only its report and publication can fail. + # A failure leaves the SUCCESS checkpoint as it is, so the retry replays the result from it + # (the branch at the top) instead of running the operator again for work that already + # finished, and reports nothing again. + try: + # Reported once the checkpoint is written, not when execute returned: the callback then + # speaks for work a retry will not run again, and a checkpoint write that fails fires + # nothing, as a plain task whose result could not be pushed fires no success callback. + # A callback that raises (a DeadlockImminentError from a synchronous SDK call in an + # async sub-task's callback) fails the sub-task without touching its checkpoint. + await self._report_item(executor, task, indexed_task_runner.report_success) + # The result is only checkpointed when the sub-task pushes XComs and returned something, + # so the same condition decides whether there is a return_value_ to push at all. + if indexed_task_state.result is not None: + await self.axcom_push(task, indexed_task_state.result) + indexed_task_runner.merge_outlet_events_into(context["outlet_events"]) + except (asyncio.CancelledError, AirflowTaskTimeout): + raise + except BaseException as e: + return task, None, e + return task, result, None + + @staticmethod + async def _record_failure( + task: IndexedTaskInstance, + runner: IndexedTaskRunner, + outcomes: IndexedTaskOutcomes, + raised: BaseException, + ) -> tuple[IndexedTaskInstance, None, BaseException]: + """Note a failed indexed task for its callbacks, checkpoint it as UP_FOR_RETRY and return it as the outcome.""" + if runner.failure is not None: + outcomes.note_failed(runner) + # Written with the input and the attempt, like the other outcomes, so the next attempt + # tells a plain retry apart from a clear or a changed input and logs only the latter. + await task.aset_state( + IndexedTaskState( + status=TaskInstanceState.UP_FOR_RETRY, + fingerprint=task.input_fingerprint, + try_number=task.try_number, + ) + ) + return task, None, raised + + @staticmethod + async def _report_item( + executor: AsyncAwareExecutor, task: IndexedTaskInstance, report: Callable[[], None] + ) -> None: + """ + Run an indexed task's success or skip report where its code ran. + + On the loop for an async operator, in a worker thread for a sync one, where a synchronous + SDK call made from the callback waits for the comms lock instead of raising. + """ + if task.is_async: + report() + else: + await executor.run_sync(report) + + @staticmethod + def _fingerprint(mapped_kwargs: Mapping[str, Any]) -> str | None: + """ + Digest one sub-task's input, stored on its checkpoint to tell whether the checkpoint still applies. + + A retry may run on another input than the attempt that wrote the checkpoints: the upstream was + cleared together with this task and produced other values. An index then no longer means the + same work, and replaying its result would hand downstream a value computed from the old value. + An input serde cannot serialize has no digest, and its checkpoint is honoured by index alone. + """ + try: + serialized = json.dumps(serialize(mapped_kwargs), sort_keys=True) + except (TypeError, ValueError, AttributeError, RecursionError): + return None + return hashlib.sha256(serialized.encode()).hexdigest() + + def _partial_inputs_from_upstream(self, unmapped_task: BaseOperator) -> dict[str, Any]: + """ + Collect the rendered values of the partial kwargs an upstream task provides. + + They belong in the fingerprint next to the iterated kwargs: clearing the upstream together with + this task can change them while the input stays the same, and a checkpoint written with the old + value must not be replayed. Only XComArg values count, read back from the unmapped operator + once rendered, at the top level or inside a mapping such as a ``@task``'s ``op_kwargs``. Other + templated values are left out on purpose: one like ``{{ ti.try_number }}`` changes with every + attempt and would make every checkpoint look stale. + """ + inputs: dict[str, Any] = {} + for key, value in self.partial_kwargs.items(): + if isinstance(value, XComArg): + inputs[key] = getattr(unmapped_task, key, None) + elif isinstance(value, Mapping): + rendered = getattr(unmapped_task, key, None) + for name, nested in value.items(): + if isinstance(nested, XComArg): + inputs[f"{key}.{name}"] = ( + rendered.get(name) if isinstance(rendered, Mapping) else None + ) + return inputs + + def _create_task( + self, + context: Context, + index: int, + mapped_kwargs: Mapping[str, Any], + jinja_env: jinja2.Environment, + ) -> IndexedTaskInstance: + unmapped_task = self._operator.unmap(mapped_kwargs) + # Make sure deferred operators will always raise a DeferredTask exception when executed + unmapped_task.start_from_trigger = False + + indexed_ti = IndexedTaskInstance.create_indexed_task( + context=context, + index=index, + operator=unmapped_task, + ) + + # Rendered against the context the indexed task will execute with (its own ti, task and + # store view), not the parent's: context_update_for_unmapped() sets ti.task in place, and a + # template must read the same task_state_store as execute() does. + self._render_unmapped_operator(indexed_ti.context_for(context), unmapped_task, jinja_env) + # Taken once rendered, so the partial kwargs an upstream provides are in it with their value. + indexed_ti.input_fingerprint = self._fingerprint( + {**mapped_kwargs, **self._partial_inputs_from_upstream(unmapped_task)} + ) + return indexed_ti + + def execute(self, context: Context): + jinja_env = self.get_template_env(dag=self.dag) + + async def tasks() -> AsyncIterator[IndexedTaskInstance]: + # Resolved and read by the executor on the running event loop, so the input's XCom + # reads go through asend and cannot deadlock with the sub-tasks' own SDK calls (see + # AsyncAwareExecutor.imap_unordered). + self._state.resolved = await self.expand_input.aresolve(context) + for index in range(self._state.resolved.length): + # Rendering may call the supervisor synchronously (an XComArg in a partial kwarg, + # ``{{ var.value.x }}``, ``{{ conn.x }}``), which raises DeadlockImminentError on the + # loop thread while a sub-task's asend is in flight. From a worker thread the same + # call waits for it instead, as in XComArg.aresolve. + yield await asyncio.to_thread( + self._create_task, + context=context, + index=index, + mapped_kwargs=await self._state.resolved.aget(index), + jinja_env=jinja_env, + ) + + try: + do_xcom_push, skipped = self._run_tasks(context=context, tasks=tasks()) + result = ( + XComIterable( + task_id=self.task_id, + dag_id=self.dag_id, + run_id=context["run_id"], + length=self._state.length, + map_index=context["ti"].map_index, + skipped=skipped, + ) + if do_xcom_push and self._state.resolved + else None + ) + if skipped: + self._skip_downstream_of_a_partial_skip(context, result) + return result + finally: + # Renewed once the run ended, not when it starts (see IterationState): a kill that lands + # before the run must stop it, and a rerun in the same process must not see that kill. + self._state = IterationState() diff --git a/task-sdk/src/airflow/sdk/definitions/mappedoperator.py b/task-sdk/src/airflow/sdk/definitions/mappedoperator.py index 3a68ecadfc22a..3121bc821cdf8 100644 --- a/task-sdk/src/airflow/sdk/definitions/mappedoperator.py +++ b/task-sdk/src/airflow/sdk/definitions/mappedoperator.py @@ -64,18 +64,32 @@ OperatorExpandArgument, OperatorExpandKwargsArgument, ) + from airflow.sdk.definitions.iterableoperator import IterableOperator from airflow.sdk.definitions.operator_resources import Resources from airflow.sdk.definitions.param import ParamsDict from airflow.sdk.definitions.retry_policy import RetryPolicy from airflow.sdk.types import WeightRuleParam from airflow.triggers.base import StartTriggerArgs -ValidationSource = Literal["expand"] | Literal["partial"] +ValidationSource = Literal["expand"] | Literal["iterate"] | Literal["partial"] + +# Raised wherever ``task_concurrency`` is given to something that is not iterated. Airflow 2 had an +# option of the same name for what is now ``max_active_tis_per_dag``, so a DAG carried over with it +# is told what to use rather than failing on an argument that looks unknown. +TASK_CONCURRENCY_REJECTED = ( + "task_concurrency is only accepted by .iterate() and .iterate_kwargs(), where it sets how many " + "iterations run at once. It is not the Airflow 2 option of that name, which is now " + "max_active_tis_per_dag." +) def validate_mapping_kwargs(op: type[BaseOperator], func: ValidationSource, value: dict[str, Any]) -> None: # use a dict so order of args is same as code order unknown_args = value.copy() + if func == "partial": + # Accepted by partial() for .iterate() to read, and rejected by .expand(), although it is not + # a BaseOperator parameter: see BaseOperator.__init__ for why it must not be one. + unknown_args.pop("task_concurrency", None) for klass in op.mro(): init = klass.__init__ # type: ignore[misc] try: @@ -84,14 +98,14 @@ def validate_mapping_kwargs(op: type[BaseOperator], func: ValidationSource, valu continue for name in param_names: value = unknown_args.pop(name, NOTSET) - if func != "expand": + if func not in ("expand", "iterate"): continue if value is NOTSET: continue if is_mappable(value): continue type_name = type(value).__name__ - error = f"{op.__name__}.expand() got an unexpected type {type_name!r} for keyword argument {name}" + error = f"{op.__name__}.{func}() got an unexpected type {type_name!r} for keyword argument {name}" raise ValueError(error) if not unknown_args: return # If we have no args left to check: stop looking at the MRO chain. @@ -196,6 +210,12 @@ def __del__(self): def expand(self, **mapped_kwargs: OperatorExpandArgument) -> MappedOperator: if not mapped_kwargs: raise TypeError("no arguments to expand against") + # task_concurrency only has meaning for Iterable Tasks (as the sub-task thread + # count consumed by IterableOperator via .iterate()/.iterate_kwargs()). + # A plain .expand() never reaches that code path, so reject it here rather than silently + # accepting a dead value. + if "task_concurrency" in self.kwargs: + raise TypeError(TASK_CONCURRENCY_REJECTED) validate_mapping_kwargs(self.operator_class, "expand", mapped_kwargs) prevent_duplicates(self.kwargs, mapped_kwargs, fail_reason="unmappable or already specified") # Since the input is already checked at parse time, we can set strict @@ -211,9 +231,18 @@ def expand_kwargs(self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") elif not isinstance(kwargs, XComArg): raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") + # See the comment in expand() above: task_concurrency has no meaning outside iterate(). + if "task_concurrency" in self.kwargs: + raise TypeError(TASK_CONCURRENCY_REJECTED) return self._expand(ListOfDictsExpandInput(kwargs), strict=strict) - def _expand(self, expand_input: ExpandInput, *, strict: bool) -> MappedOperator: + def _expand( + self, + expand_input: ExpandInput, + *, + strict: bool, + register_with_dag: bool = True, + ) -> MappedOperator: from airflow.providers.standard.operators.empty import EmptyOperator from airflow.sdk import BaseSensorOperator from airflow.sdk.bases.skipmixin import SkipMixin @@ -243,7 +272,7 @@ def _expand(self, expand_input: ExpandInput, *, strict: bool) -> MappedOperator: except AttributeError: operator_name = self.operator_class.__name__ - op = MappedOperator( + return MappedOperator( operator_class=self.operator_class, expand_input=expand_input, partial_kwargs=partial_kwargs, @@ -273,8 +302,47 @@ def _expand(self, expand_input: ExpandInput, *, strict: bool) -> MappedOperator: # TODO: Move these to task SDK's BaseOperator and remove getattr start_trigger_args=start_trigger_args, start_from_trigger=start_from_trigger, + register_with_dag=register_with_dag, ) - return op + + def iterate(self, **mapped_kwargs: OperatorExpandArgument) -> IterableOperator: + """ + Iterate the operator over ``mapped_kwargs`` inside a single task instance. + + The counterpart of :meth:`expand` for Iterable Tasks: the same inputs, but processed by one + :class:`~airflow.sdk.definitions.iterableoperator.IterableOperator` instead of one task + instance per item. + """ + if not mapped_kwargs: + raise TypeError("no arguments to iterate against") + + validate_mapping_kwargs(self.operator_class, "iterate", mapped_kwargs) + prevent_duplicates(self.kwargs, mapped_kwargs, fail_reason="unmappable or already specified") + # Since the input is already checked at parse time, we can set strict + # to False to skip the checks on execution. + return self._iterate(DictOfListsExpandInput(mapped_kwargs), strict=False) + + def iterate_kwargs( + self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = True + ) -> IterableOperator: + """Iterate the operator over a list of dicts or an XComArg; see :meth:`iterate`.""" + from airflow.sdk.definitions.xcom_arg import XComArg + + if isinstance(kwargs, Sequence): + for item in kwargs: + if not isinstance(item, (XComArg, Mapping)): + raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") + elif not isinstance(kwargs, XComArg): + raise TypeError(f"expected XComArg or list[dict], not {type(kwargs).__name__}") + return self._iterate(ListOfDictsExpandInput(kwargs), strict=strict) + + def _iterate(self, expand_input: ExpandInput, *, strict: bool) -> IterableOperator: + from airflow.sdk.definitions.iterableoperator import IterableOperator + + # The MappedOperator only drives the iteration in memory: it is never registered with the + # DAG, the IterableOperator is the single real task. _expand marks the partial as consumed. + operator = self._expand(expand_input, strict=strict, register_with_dag=False) + return IterableOperator(operator=operator, expand_input=expand_input) @attrs.define( @@ -325,6 +393,7 @@ class MappedOperator(AbstractOperator): end_date: pendulum.DateTime | None upstream_task_ids: set[str] = attrs.field(factory=set, init=False) downstream_task_ids: set[str] = attrs.field(factory=set, init=False) + _register_with_dag: bool = attrs.field(alias="register_with_dag", default=True) _disallow_kwargs_override: bool """Whether execution fails if ``expand_input`` has duplicates to ``partial_kwargs``. @@ -346,19 +415,26 @@ def __repr__(self): return f"" def __attrs_post_init__(self): - from airflow.sdk.definitions.xcom_arg import XComArg - - if self.get_closest_mapped_task_group() is not None: - raise NotImplementedError("operator expansion in an expanded task group is not yet supported") - - if self.task_group: - self.task_group.add(self) - if self.dag: - self.dag.add_task(self) - XComArg.apply_upstream_relationship(self, self._get_specified_expand_input().value) - for k, v in self.partial_kwargs.items(): - if k in self.template_fields: - XComArg.apply_upstream_relationship(self, v) + # When _register_with_dag is False (i.e. IterableOperator), we intentionally + # skip the *entire* body — not just XComArg.apply_upstream_relationship. + # IterableOperator creates in-memory MappedOperator instances solely to drive task + # iteration; they must NOT be registered with the DAG or task group because Airflow + # treats the IterableOperator itself as the single real task instance in the DB. + # Calling dag.add_task() or task_group.add() here would raise duplicate-task errors. + if self._register_with_dag: + from airflow.sdk.definitions.xcom_arg import XComArg + + if self.get_closest_mapped_task_group() is not None: + raise NotImplementedError("operator expansion in an expanded task group is not yet supported") + + if self.task_group: + self.task_group.add(self) + if self.dag: + self.dag.add_task(self) + XComArg.apply_upstream_relationship(self, self._get_specified_expand_input().value) + for k, v in self.partial_kwargs.items(): + if k in self.template_fields: + XComArg.apply_upstream_relationship(self, v) @methodtools.lru_cache(maxsize=None) @classmethod @@ -367,6 +443,7 @@ def get_serialized_fields(cls): return frozenset(attrs.fields_dict(MappedOperator)) - { "_is_empty", "_can_skip_downstream", + "_register_with_dag", "dag", "deps", "expand_input", # This is needed to be able to accept XComArg. @@ -792,6 +869,9 @@ def unmap(self, resolve: Mapping[str, Any]) -> BaseOperator: is_setup = kwargs.pop("is_setup", False) is_teardown = kwargs.pop("is_teardown", False) on_failure_fail_dagrun = kwargs.pop("on_failure_fail_dagrun", False) + # task_concurrency is iterable-task metadata (the sub-task worker count), not an operator + # init argument. + kwargs.pop("task_concurrency", None) kwargs["task_id"] = self.task_id op = self.operator_class(**kwargs, _airflow_from_mapped=True) op.is_setup = is_setup diff --git a/task-sdk/src/airflow/sdk/definitions/xcom_arg.py b/task-sdk/src/airflow/sdk/definitions/xcom_arg.py index a38e47a466bcf..0631e68131a7a 100644 --- a/task-sdk/src/airflow/sdk/definitions/xcom_arg.py +++ b/task-sdk/src/airflow/sdk/definitions/xcom_arg.py @@ -17,6 +17,7 @@ from __future__ import annotations +import asyncio import contextlib import inspect import itertools @@ -78,6 +79,11 @@ class XComArg(ResolveMixin, DependencyMixin): :param operator: Operator instance to which the XComArg references. :param key: Key used to pull the XCom value. Defaults to *XCOM_RETURN_KEY*, i.e. the referenced operator's return value. + + **Subclassing.** A subclass has to implement :meth:`iter_references` and :meth:`resolve`. The + one other method an iterated task (``.iterate()``) uses, :meth:`aresolve`, has a working + default built on those two: it runs ``resolve`` in a worker thread. Override it only to do + better, as the built-in subclasses do by pulling through ``ti.axcom_pull`` directly. """ @overload @@ -180,6 +186,19 @@ def concat(self, *others: XComArg) -> ConcatXComArg: def resolve(self, context: Mapping[str, Any]) -> Any: raise NotImplementedError() + async def aresolve(self, context: Mapping[str, Any]) -> Any: + """ + Async twin of :meth:`resolve`, for callers running on the task's event loop. + + The default runs :meth:`resolve` in a worker thread, so any subclass that implements + ``resolve`` works on the iterated path unchanged: a blocking supervisor call made from that + thread waits for the in-flight ``asend`` calls of other sub-tasks instead of deadlocking + with them, which the same call on the loop thread would do (see + ``AsyncAwareExecutor.imap_unordered``). The built-in subclasses override this to pull + through ``ti.axcom_pull`` directly and skip the thread hand-off. + """ + return await asyncio.to_thread(self.resolve, context) + def __enter__(self): if not self.operator.is_setup and not self.operator.is_teardown: raise AirflowException("Only setup/teardown tasks can be used as context managers.") @@ -331,6 +350,40 @@ def concat(self, *others: XComArg) -> ConcatXComArg: def resolve(self, context: Mapping[str, Any]) -> Any: ti = context["ti"] + map_indexes = self._resolve_map_indexes(ti) + if isinstance(map_indexes, LazyXComSequence): + return map_indexes + result = ti.xcom_pull( + task_ids=self.operator.task_id, + key=self.key, + default=NOTSET, + map_indexes=map_indexes, + ) + return self._check_pulled(ti, result) + + async def aresolve(self, context: Mapping[str, Any]) -> Any: + ti = context["ti"] + # Computing the map indexes may count upstream task instances through a synchronous + # supervisor call, so keep it off the loop thread. + map_indexes = await asyncio.to_thread(self._resolve_map_indexes, ti) + if isinstance(map_indexes, LazyXComSequence): + return map_indexes + result = await ti.axcom_pull( + task_ids=self.operator.task_id, + key=self.key, + default=NOTSET, + map_indexes=map_indexes, + ) + return self._check_pulled(ti, result) + + def _resolve_map_indexes(self, ti: Any) -> LazyXComSequence | int | range | None: + """ + Return the upstream map indexes to pull, or a lazy sequence over the whole mapped upstream. + + A ``LazyXComSequence`` is returned when the upstream is mapped (or expands to a whole mapped + task group), so nothing is pulled eagerly; otherwise ``None`` pulls the unmapped instance and + an int or range pulls the relevant expanded instances. + """ task_id = self.operator.task_id if self.operator.is_mapped: @@ -360,12 +413,10 @@ def resolve(self, context: Mapping[str, Any]) -> Any: # the value is actually returned and pushed to XCom (see airflow.sdk.serde.serialize). return LazyXComSequence(xcom_arg=self, ti=ti) map_indexes = computed - result = ti.xcom_pull( - task_ids=task_id, - key=self.key, - default=NOTSET, - map_indexes=map_indexes, - ) + return map_indexes + + def _check_pulled(self, ti: Any, result: Any) -> Any: + """Turn a pulled XCom into the resolved value, or raise when the XCom must exist.""" if is_arg_set(result): return result if self.key == BaseXCom.XCOM_RETURN_KEY: @@ -377,7 +428,7 @@ def resolve(self, context: Mapping[str, Any]) -> Any: # different names than the predefined "XCOM_RETURN_KEY" and won't be found. # Therefore, it's better to return "None" like we did above where self.key==XCOM_RETURN_KEY. return None - raise XComNotFound(ti.dag_id, task_id, self.key) + raise XComNotFound(ti.dag_id, self.operator.task_id, self.key) def _get_callable_name(f: Callable | str) -> str: @@ -448,7 +499,12 @@ def map(self, f: Callable[[Any], Any]) -> MapXComArg: return MapXComArg(self.arg, [*self.callables, f]) def resolve(self, context: Mapping[str, Any]) -> Any: - value = self.arg.resolve(context) + return self._map(self.arg.resolve(context)) + + async def aresolve(self, context: Mapping[str, Any]) -> Any: + return self._map(await self.arg.aresolve(context)) + + def _map(self, value: Any) -> _MapResult: if not isinstance(value, (Sequence, dict)): raise ValueError(f"XCom map expects sequence or dict, not {type(value).__name__}") return _MapResult(value, self.callables) @@ -510,7 +566,12 @@ def iter_references(self) -> Iterator[tuple[Operator, str]]: yield from arg.iter_references() def resolve(self, context: Mapping[str, Any]) -> Any: - values = [arg.resolve(context) for arg in self.args] + return self._zip([arg.resolve(context) for arg in self.args]) + + async def aresolve(self, context: Mapping[str, Any]) -> Any: + return self._zip([await arg.aresolve(context) for arg in self.args]) + + def _zip(self, values: list[Any]) -> _ZipResult: for value in values: if not isinstance(value, (Sequence, dict)): raise ValueError(f"XCom zip expects sequence or dict, not {type(value).__name__}") @@ -571,7 +632,12 @@ def concat(self, *others: XComArg) -> ConcatXComArg: return ConcatXComArg([*self.args, *others]) def resolve(self, context: Mapping[str, Any]) -> Any: - values = [arg.resolve(context) for arg in self.args] + return self._concat([arg.resolve(context) for arg in self.args]) + + async def aresolve(self, context: Mapping[str, Any]) -> Any: + return self._concat([await arg.aresolve(context) for arg in self.args]) + + def _concat(self, values: list[Any]) -> _ConcatResult: for value in values: if not isinstance(value, (Sequence, dict)): raise ValueError(f"XCom concat expects sequence or dict, not {type(value).__name__}") diff --git a/task-sdk/src/airflow/sdk/execution_time/comms.py b/task-sdk/src/airflow/sdk/execution_time/comms.py index 6472863fbd8a7..45e5e8bde4a5f 100644 --- a/task-sdk/src/airflow/sdk/execution_time/comms.py +++ b/task-sdk/src/airflow/sdk/execution_time/comms.py @@ -308,6 +308,38 @@ def send(self, msg: SendMsgType) -> ReceiveMsgType | None: finally: self._thread_lock.release() + async def _acquire_thread_lock(self, loop: asyncio.AbstractEventLoop) -> None: + """ + Take ``_thread_lock`` in a worker thread, so the loop keeps running, and keep it balanced. + + Cancelling the wait cancels the future, not the thread: the acquire completes later and the + lock would stay taken with nobody to release it, so every later ``send`` would block for + good. The hand-off settles who releases it: the thread, when it finds the wait abandoned, + or the caller's ``finally`` when the acquire had completed before the cancellation landed. + """ + handoff = threading.Lock() + abandoned = False + acquired = False + + def acquire() -> None: + nonlocal acquired + self._thread_lock.acquire() + with handoff: + if abandoned: + self._thread_lock.release() + return + acquired = True + + try: + await loop.run_in_executor(None, acquire) + except BaseException: + with handoff: + if acquired: + self._thread_lock.release() + else: + abandoned = True + raise + async def asend(self, msg: SendMsgType) -> ReceiveMsgType | None: """ Send a request to the parent without blocking. @@ -321,7 +353,7 @@ async def asend(self, msg: SendMsgType) -> ReceiveMsgType | None: async with self._async_lock: # Acquire the threading lock without blocking the event loop loop = asyncio.get_running_loop() - await loop.run_in_executor(None, self._thread_lock.acquire) + await self._acquire_thread_lock(loop) try: # Async write to socket await loop.sock_sendall(self.socket, frame_bytes) diff --git a/task-sdk/src/airflow/sdk/execution_time/context.py b/task-sdk/src/airflow/sdk/execution_time/context.py index 02be651e9e2ae..892100c560f3a 100644 --- a/task-sdk/src/airflow/sdk/execution_time/context.py +++ b/task-sdk/src/airflow/sdk/execution_time/context.py @@ -34,7 +34,7 @@ from airflow.sdk._shared.state import AssetScope from airflow.sdk.configuration import conf -from airflow.sdk.definitions._internal.contextmanager import _CURRENT_CONTEXT +from airflow.sdk.definitions._internal.contextmanager import _CURRENT_CONTEXT, _INDEXED_CONTEXT from airflow.sdk.definitions._internal.types import NOTSET from airflow.sdk.definitions.asset import ( Asset, @@ -148,7 +148,6 @@ }, } - log = structlog.get_logger(logger_name="task") #: Pass as ``retention`` to ``task_state_store.set()`` to store a key that never expires, @@ -885,6 +884,80 @@ def _clear_backend_only(self) -> None: backend.clear(self._scope) +class IndexedTaskStateStoreAccessor(TaskStateStoreAccessor): + """ + The parent task instance's state store as seen from one iteration of an iterated task. + + Every iteration runs under the same task instance, so a key written from inside one would be + overwritten by its siblings. This view suffixes each key with the iteration's index, exactly as + ``IndexedTaskInstance.xcom_push`` does for XComs, so ``task_state_store.set("last_offset", 3)`` + in iteration 2 lands under ``last_offset_2``. Clearing is refused: it would wipe the siblings' + state and the operator's own checkpoints; an iteration deletes its own keys instead. Keys + starting with ``_iterable`` are refused as well: the operator keeps its checkpoints in this + store under ``_iterable_``, which is what ``_iterable`` would become once suffixed. + + The view wraps the parent's accessor instead of being one: it has no ``_ti_id`` or ``_scope`` + of its own, so every inherited member that reads them is overridden here. + """ + + # The namespace of IndexedTaskState.build_key and Checkpoints.COMPLETION_KEY. + RESERVED_PREFIX = "_iterable" + + def __init__(self, store: TaskStateStoreAccessor, index: int) -> None: + self._store = store + self._index = index + + def __eq__(self, other: object) -> bool: + if not isinstance(other, IndexedTaskStateStoreAccessor): + return False + return self._store == other._store and self._index == other._index + + def __hash__(self) -> int: + return hash((self._store, self._index)) + + def __repr__(self) -> str: + return f"" + + def _indexed(self, key: str) -> str: + if key.startswith(self.RESERVED_PREFIX): + raise ValueError( + f"task_state_store keys starting with {self.RESERVED_PREFIX!r} are reserved for the " + f"checkpoints of the iterated task, got {key!r}" + ) + return f"{key}_{self._index}" + + def get(self, key: str, default: JsonValue = None) -> JsonValue: + return self._store.get(self._indexed(key), default) + + async def aget(self, key: str, default: JsonValue = None) -> JsonValue: + return await self._store.aget(self._indexed(key), default) + + def set(self, key: str, value: JsonValue, *, retention: timedelta | None = None) -> None: + self._store.set(self._indexed(key), value, retention=retention) + + async def aset(self, key: str, value: JsonValue, *, retention: timedelta | None = None) -> None: + await self._store.aset(self._indexed(key), value, retention=retention) + + def delete(self, key: str) -> None: + self._store.delete(self._indexed(key)) + + async def adelete(self, key: str) -> None: + await self._store.adelete(self._indexed(key)) + + def clear(self) -> None: + raise RuntimeError( + "task_state_store.clear() is not available inside an iterated task: the store is shared " + "with the other iterations and the operator's checkpoints; delete your own keys instead" + ) + + def _clear_backend_only(self) -> None: + # The runner's clear_on_success path, on the parent's accessor; refused here as clear() is. + self.clear() + + async def aclear(self) -> None: + self.clear() + + class AssetStateStoreAccessor: """ Accessor for asset store scoped to a single asset. @@ -1356,9 +1429,10 @@ def __getitem__(self, key: Asset | AssetAlias | AssetRef) -> OutletEventAccessor else: raise TypeError(f"Key should be either an asset or an asset alias, not {type(key)}") - if hashable_key not in self._dict: - self._dict[hashable_key] = OutletEventAccessor(extra={}, key=hashable_key) - return self._dict[hashable_key] + # setdefault is atomic under the GIL: if two threads race on the same + # key the first writer wins and both threads get back the same accessor, + # so neither thread's accumulated events are silently discarded. + return self._dict.setdefault(hashable_key, OutletEventAccessor(extra={}, key=hashable_key)) @attrs.define(init=False) @@ -1619,6 +1693,22 @@ def set_current_context(context: Context) -> Generator[Context, None, None]: ) +@contextlib.contextmanager +def set_indexed_context(context: Context) -> Generator[Context, None, None]: + """ + Make ``context`` the current context of one iteration of an iterated task, for this block. + + Seen only by the thread or asyncio task that entered the block, so iterations running + concurrently never see each other's context. Anything else, a thread the iteration starts + included, sees the task's own context set by :func:`set_current_context`. + """ + token = _INDEXED_CONTEXT.set((*_INDEXED_CONTEXT.get(), context)) + try: + yield context + finally: + _INDEXED_CONTEXT.reset(token) + + def context_update_for_unmapped(context: Context, task: BaseOperator) -> None: """ Update context after task unmapping. diff --git a/task-sdk/src/airflow/sdk/execution_time/executor.py b/task-sdk/src/airflow/sdk/execution_time/executor.py new file mode 100644 index 0000000000000..c1b50dbe657f6 --- /dev/null +++ b/task-sdk/src/airflow/sdk/execution_time/executor.py @@ -0,0 +1,268 @@ +# +# 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 + +import inspect +import logging +import os +import time +from asyncio import ( + FIRST_COMPLETED, + AbstractEventLoop, + CancelledError, + Future, + Semaphore, + Task, + TimeoutError as AsyncTimeoutError, + gather, + wait, + wait_for, + wrap_future, +) +from collections.abc import AsyncIterable, Callable, Iterator +from concurrent.futures import Executor, ThreadPoolExecutor +from contextlib import suppress +from typing import Any + +_log = logging.getLogger(__name__) + + +class AsyncAwareExecutor(Executor): + """ + Executes both sync and async functions concurrently. + + Sync functions run in a ThreadPoolExecutor. + Async coroutines run on an asyncio event loop with a semaphore limit. + + :param loop: Event loop used to schedule async tasks and coordinate mixed execution. + :param max_workers: Maximum concurrent workers used by both thread pool and async semaphore. + :param shutdown_timeout: Maximum time to wait, in seconds, for in-flight async tasks and + thread-pool workers to finish during ``shutdown(wait=True)``. Python threads cannot be + forcibly stopped, so a worker stuck in blocking user code (slow HTTP call, blocked C + extension, a deadlocked DB driver, ...) would otherwise hang ``shutdown()`` forever. + """ + + def __init__( + self, loop: AbstractEventLoop, max_workers: int | None = None, shutdown_timeout: float = 10.0 + ): + if max_workers is None: + max_workers = os.cpu_count() or 1 + if max_workers <= 0: + raise ValueError("max_workers must be greater than 0") + + self._loop = loop + self._max_workers = max_workers + self._shutdown_timeout = shutdown_timeout + self._semaphore = Semaphore(max_workers) + self._thread_pool = ThreadPoolExecutor(max_workers=max_workers) + self._async_tasks: set[Task[Any]] = set() + self._shutdown = False + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + if exc_type is not None: + # On error path, cancel futures but still wait briefly for them to + # process CancelledError and release resources (e.g., threading + # locks). Without waiting, cancelled tasks that hold _thread_lock + # never execute their finally blocks, permanently leaking the lock + # and causing subsequent comms.send() calls to deadlock. A task + # cancelled while still waiting for that lock is CommsDecoder's + # concern: its acquire goes on in a thread, and asend releases what + # that thread takes after the wait was abandoned. + self.shutdown(wait=True, cancel_futures=True) + else: + self.shutdown(wait=True) + + def shutdown(self, wait: bool = True, *, cancel_futures: bool = False) -> None: + if self._shutdown: + return + + self._shutdown = True + + if cancel_futures: + for task in list(self._async_tasks): + task.cancel() + + if wait and self._async_tasks: + with suppress(TimeoutError, AsyncTimeoutError): + self._loop.run_until_complete( + wait_for( + gather(*self._async_tasks, return_exceptions=True), + timeout=self._shutdown_timeout, + ) + ) + + # ThreadPoolExecutor.shutdown(wait=True) blocks until every worker thread + # finishes its current work item, with no way to bound that wait or forcibly + # stop a thread stuck in blocking user code (slow HTTP call, blocked C + # extension, a deadlocked DB driver, ...). Ask the pool to stop accepting new + # work (and cancel anything not yet started) up front, then bound how long we + # personally wait on the worker threads instead of blocking indefinitely. + self._thread_pool.shutdown(wait=False, cancel_futures=cancel_futures) + + if wait: + threads = list(getattr(self._thread_pool, "_threads", ())) + deadline = time.monotonic() + self._shutdown_timeout + for thread in threads: + remaining = deadline - time.monotonic() + if remaining <= 0: + break + thread.join(timeout=remaining) + + stuck = [thread.name for thread in threads if thread.is_alive()] + if stuck: + _log.error( + "%d worker thread(s) still running %.1fs after shutdown was requested; " + "giving up waiting to avoid blocking indefinitely. This may leak resources. " + "Affected threads: %s", + len(stuck), + self._shutdown_timeout, + stuck, + ) + + def submit(self, func: Callable[..., Any] | Any, *args, **kwargs) -> Future[Any]: # type: ignore[override] + """ + Submit a callable for execution. + + Always returns an asyncio.Future for consistency, whether the callable + is sync (run in thread pool) or async (run on the event loop). + """ + if self._shutdown: + raise RuntimeError("cannot schedule new futures after shutdown") + + if inspect.iscoroutine(func): + coro = func + elif inspect.iscoroutinefunction(func): + coro = func(*args, **kwargs) + else: + # Wrap thread pool future as asyncio.Future for consistent return type + return wrap_future(self._thread_pool.submit(func, *args, **kwargs), loop=self._loop) + + async def guarded(): + try: + async with self._semaphore: + return await coro + except CancelledError: + # If cancellation occurs while waiting for the semaphore, + # the inner coroutine was never awaited. Close it to prevent + # "coroutine was never awaited" RuntimeWarning. + coro.close() + raise + + task = self._loop.create_task(guarded()) + self._async_tasks.add(task) + task.add_done_callback(self._async_tasks.discard) + return task + + async def run_sync(self, func: Callable[..., Any], *args, **kwargs) -> Any: + """Run a sync callable in this executor's thread pool and await its result.""" + future = self._thread_pool.submit(func, *args, **kwargs) + return await wrap_future(future, loop=self._loop) + + def imap_unordered( + self, + fn: Callable[..., Any], + *iterables: AsyncIterable[Any], + timeout: float | None = None, + stop: Callable[[], bool] | None = None, + ) -> Iterator[Any]: + """ + Apply ``fn`` to async iterables, zipped, and stream results in completion order. + + Named after ``multiprocessing.Pool.imap_unordered`` because that is the contract: results + come back as calls finish, not in submission order. It deliberately does not override + ``concurrent.futures.Executor.map``, which promises submission order, so code holding a + plain ``Executor`` keeps that guarantee and this method has to be asked for by name. + + The iterables are async, unlike ``Executor.map``'s: items are pulled and + calls submitted from a coroutine on the running loop, never from the main thread between + two ``run_until_complete`` calls. At such a moment a call can be parked mid-``asend`` + holding the supervisor channel's thread lock; a synchronous SDK call pulling the next item + (an XCom read behind an iterated task's input) would then take that lock in blocking mode, + since no loop is running, while the holder needs the loop to run to release it, and the + process freezes with every thread idle. On the running loop an async iterable reads + through ``asend``, and a synchronous SDK call from the loop thread meets the SDK's + running-loop check and raises ``DeadlockImminentError`` instead of hanging. + + Results are handed to the caller while the loop is paused, which is safe: the caller only + consumes them. + + ``stop`` is asked before and after every pull; once it answers True nothing more is + submitted, an item pulled at that moment included, and the calls already submitted are + drained as usual. A kill sets such a flag, + so that killing what is in flight is not followed by starting the next items. + """ + if self._shutdown: + raise RuntimeError("cannot schedule new futures after shutdown") + + start = time.monotonic() + iterators = [iterable.__aiter__() for iterable in iterables] + pending: set[Future[Any]] = set() + exhausted = False + + def _remaining_timeout() -> float | None: + if timeout is None: + return None + remaining = timeout - (time.monotonic() - start) + if remaining <= 0: + raise TimeoutError() + return remaining + + async def _next_args() -> tuple[Any, ...] | None: + """Return the next argument tuple, or None once any iterable is exhausted (like ``zip``).""" + args = [] + for iterator in iterators: + try: + args.append(await iterator.__anext__()) + except StopAsyncIteration: + return None + return tuple(args) + + async def _fill_pending() -> None: + """Submit calls until pending reaches max_workers or an iterable is exhausted.""" + nonlocal exhausted + while not exhausted and len(pending) < self._max_workers: + if stop is not None and stop(): + exhausted = True + return + args = await _next_args() + if args is None or (stop is not None and stop()): + # Exhausted, or stopped while the item was pulled: an item pulled at that + # moment is dropped rather than started. + exhausted = True + return + pending.add(self.submit(fn, *args)) + + # Every pull from the iterables runs on the loop, the initial one included. + self._loop.run_until_complete(_fill_pending()) + + while pending: + done, _ = self._loop.run_until_complete( + wait(pending, timeout=_remaining_timeout(), return_when=FIRST_COMPLETED) + ) + + if not done: + raise TimeoutError() + + for completed in done: + pending.discard(completed) + yield completed.result() + + self._loop.run_until_complete(_fill_pending()) diff --git a/task-sdk/src/airflow/sdk/execution_time/lazy_sequence.py b/task-sdk/src/airflow/sdk/execution_time/lazy_sequence.py index 4efb0b71368ca..2344a48591721 100644 --- a/task-sdk/src/airflow/sdk/execution_time/lazy_sequence.py +++ b/task-sdk/src/airflow/sdk/execution_time/lazy_sequence.py @@ -19,7 +19,7 @@ import collections import itertools -from collections.abc import Iterator, Sequence +from collections.abc import AsyncIterator, Iterator, Sequence from typing import TYPE_CHECKING, Any, Literal, TypeVar, overload import attrs @@ -60,6 +60,31 @@ def __iter__(self) -> Iterator[T]: return self +@attrs.define +class AsyncLazyXComIterator(AsyncIterator[T]): + """ + Async twin of :class:`LazyXComIterator`: the same item reads, sent through ``asend``. + + An iterated task consumes a mapped task's results as its input on the event loop; reading + them synchronously there would block the loop thread on the supervisor channel while the + sub-tasks' own ``asend`` calls are in flight (see ``AsyncAwareExecutor.imap_unordered``). + """ + + seq: LazyXComSequence[T] + index: int = 0 + + def __aiter__(self) -> AsyncIterator[T]: + return self + + async def __anext__(self) -> T: + try: + val = await self.seq.aget(self.index) + except IndexError: + raise StopAsyncIteration from None + self.index += 1 + return val + + @attrs.define class LazyXComSequence(Sequence[T]): _len: int | None = attrs.field(init=False, default=None) @@ -87,6 +112,35 @@ def __hash__(self): def __iter__(self) -> Iterator[T]: return LazyXComIterator(seq=self) + def __aiter__(self) -> AsyncIterator[T]: + return AsyncLazyXComIterator(seq=self) + + async def aget(self, index: int) -> T: + """Async counterpart of ``self[index]``; the same ``GetXComSequenceItem`` request via ``asend``.""" + from airflow.sdk.execution_time.comms import ( + ErrorResponse, + GetXComSequenceItem, + XComSequenceIndexResult, + ) + from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS + from airflow.sdk.execution_time.xcom import XCom + + source = (xcom_arg := self._xcom_arg).operator + msg = await SUPERVISOR_COMMS.asend( + GetXComSequenceItem( + key=xcom_arg.key, + dag_id=source.dag_id, + task_id=source.task_id, + run_id=self._ti.run_id, + offset=index, + ), + ) + if isinstance(msg, ErrorResponse): + raise IndexError(index) + if not isinstance(msg, XComSequenceIndexResult): + raise TypeError(f"Got unexpected response to GetXComSequenceItem: {msg!r}") + return XCom.deserialize_value(_XComWrapper(msg.root)) + def __len__(self) -> int: if self._len is None: from airflow.sdk.execution_time.comms import ErrorResponse, GetXComCount, XComCountResponse @@ -109,6 +163,29 @@ def __len__(self) -> int: self._len = msg.len return self._len + async def alen(self) -> int: + """Async twin of ``len(self)``: the same ``GetXComCount`` request, sent with ``asend``.""" + if self._len is None: + from airflow.sdk.execution_time.comms import ErrorResponse, GetXComCount, XComCountResponse + from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS + + task = self._xcom_arg.operator + + msg = await SUPERVISOR_COMMS.asend( + GetXComCount( + key=self._xcom_arg.key, + dag_id=task.dag_id, + run_id=self._ti.run_id, + task_id=task.task_id, + ), + ) + if isinstance(msg, ErrorResponse): + raise RuntimeError(msg) + if not isinstance(msg, XComCountResponse): + raise TypeError(f"Got unexpected response to GetXComCount: {msg!r}") + self._len = msg.len + return self._len + @overload def __getitem__(self, key: int) -> T: ... diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 2be985c04249d..95f2c8c153d67 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -2473,6 +2473,16 @@ def send(self, msg: BaseModel): return self._get_response() + async def asend(self, msg: BaseModel): + """ + Send a request to the supervisor without blocking the event loop. + + ``_handle_request`` is synchronous and in-process (no actual socket I/O), so this simply + mirrors :meth:`send` under an ``async def`` for callers (e.g. IterableOperator's checkpoint + reads/writes via the Task State Store) that require an awaitable ``asend``. + """ + return self.send(msg) + @attrs.define class TaskRunResult: diff --git a/task-sdk/src/airflow/sdk/execution_time/task_runner.py b/task-sdk/src/airflow/sdk/execution_time/task_runner.py index 9c4a999c27a2b..74cf65d10f7c3 100644 --- a/task-sdk/src/airflow/sdk/execution_time/task_runner.py +++ b/task-sdk/src/airflow/sdk/execution_time/task_runner.py @@ -22,16 +22,20 @@ import contextvars import functools import inspect +import logging import os import sys +import threading import time +from asyncio import CancelledError from collections.abc import Callable, Iterable, Iterator, Mapping from contextlib import ExitStack, contextmanager, suppress -from dataclasses import replace +from dataclasses import dataclass, replace from datetime import UTC, datetime, timedelta +from functools import cached_property from itertools import product from pathlib import Path -from typing import TYPE_CHECKING, Annotated, Any, Literal, NoReturn, cast +from typing import TYPE_CHECKING, Annotated, Any, Literal, NoReturn, Protocol, cast from urllib.parse import quote import attrs @@ -56,20 +60,23 @@ TaskInstanceState, TIRunContext, ) -from airflow.sdk.bases.operator import BaseOperator, ExecutorSafeguard +from airflow.sdk.bases.operator import BaseAsyncOperator, BaseOperator, ExecutorSafeguard from airflow.sdk.bases.skipmixin import XCOM_SKIPMIXIN_KEY from airflow.sdk.bases.xcom import BaseXCom from airflow.sdk.configuration import conf from airflow.sdk.coordinators._dag_importer import find_claiming_importer from airflow.sdk.definitions._internal.dag_parsing_context import _airflow_parsing_context_manager +from airflow.sdk.definitions._internal.logging_mixin import LoggingMixin from airflow.sdk.definitions._internal.types import NOTSET, ArgNotSet, is_arg_set from airflow.sdk.definitions.asset import ( Asset, AssetAlias, + AssetAliasEvent, AssetNameRef, AssetUniqueKey, AssetUriRef, ) +from airflow.sdk.definitions.context import clone_context from airflow.sdk.definitions.mappedoperator import MappedOperator from airflow.sdk.definitions.param import process_params from airflow.sdk.exceptions import ( @@ -79,6 +86,7 @@ AirflowRescheduleException, AirflowRuntimeError, AirflowSensorTimeout, + AirflowSkipException, AirflowTaskTerminated, AirflowTaskTimeout, ErrorType, @@ -132,6 +140,7 @@ from airflow.sdk.execution_time.context import ( AssetStateStoreAccessors, ConnectionAccessor, + IndexedTaskStateStoreAccessor, InletEventsAccessors, MacrosAccessor, OutletEventAccessors, @@ -143,6 +152,7 @@ context_to_airflow_vars, get_previous_dagrun_success, set_current_context, + set_indexed_context, ) from airflow.sdk.execution_time.email_backend import ( _DEFAULT_EMAIL_BACKEND, @@ -154,7 +164,12 @@ from airflow.sdk.execution_time.xcom import XCom from airflow.sdk.listener import get_listener_manager from airflow.sdk.observability.metrics import stats_utils -from airflow.sdk.serde import allow_class, iter_pydantic_models +from airflow.sdk.serde import ( + allow_class, + deserialize as serde_deserialize, + iter_pydantic_models, + serialize as serde_serialize, +) from airflow.sdk.state import TaskScope from airflow.sdk.timezone import coerce_datetime @@ -279,6 +294,26 @@ def __rich_repr__(self): __rich_repr__.angular = True # type: ignore[attr-defined] + @cached_property + def logical_date(self) -> datetime | None: + if self._ti_context_from_server: + dag_run = self._ti_context_from_server.dag_run + + return dag_run.logical_date + return None + + @cached_property + def task_state_store(self) -> TaskStateStoreAccessor: + return TaskStateStoreAccessor( + ti_id=self.id, + scope=TaskScope( + dag_id=self.dag_id, + run_id=self.run_id, + task_id=self.task_id, + map_index=self.map_index if self.map_index is not None else -1, + ), + ) + @detail_span("get_template_context") def get_template_context(self) -> Context: # TODO: Move this to `airflow.sdk.execution_time.context` @@ -321,15 +356,7 @@ def get_template_context(self) -> Context: "value": VariableAccessor(deserialize_json=False), }, "conn": ConnectionAccessor(), - "task_state_store": TaskStateStoreAccessor( - ti_id=self.id, - scope=TaskScope( - dag_id=self.dag_id, - run_id=self.run_id, - task_id=self.task_id, - map_index=self.map_index if self.map_index is not None else -1, - ), - ), + "task_state_store": self.task_state_store, } _asset_types = (Asset, AssetNameRef, AssetUriRef, AssetAlias) if any(isinstance(i, _asset_types) for i in self.task.inlets + self.task.outlets): @@ -363,7 +390,7 @@ def get_template_context(self) -> Context: } self._cached_template_context.update(context_from_server) - if logical_date := coerce_datetime(dag_run.logical_date): + if logical_date := coerce_datetime(self.logical_date): if TYPE_CHECKING: assert isinstance(logical_date, DateTime) ds = logical_date.strftime("%Y-%m-%d") @@ -892,6 +919,579 @@ def mark_success_url(self) -> str: return self.log_url +@dataclass +class IndexedTaskState: + status: TaskInstanceState + result: Any | None = None + # Outlet asset events the sub-task recorded on a previous successful attempt. Outlet events + # only reach the server on the parent's success payload (_handle_current_task_success), so an + # attempt that fails never registers what its succeeded sub-tasks emitted. A sub-task skipped + # on retry (because it already succeeded) never re-executes, so it never re-emits into the + # fresh OutletEventAccessors created for the new attempt; persisting a snapshot here lets + # IterableOperator._run_task replay it, which is the only way those events survive, and it + # cannot double-emit because the failed attempt sent nothing. + outlet_events: list[dict[str, Any]] | None = None + # Digest of the input the sub-task ran with. An index only means the same work while its input + # is the same, so a checkpoint is honoured on a later attempt only when this still matches: + # after the upstream was cleared and produced other values, the sub-task runs again. + fingerprint: str | None = None + # The attempt that wrote the checkpoint. After a manual clear only the checkpoints written since + # are resumed from (see Checkpoints), which this tells apart from those left by the run before. + try_number: int = 0 + # The other XComs the sub-task pushed, by their key before the index suffix. The runner deletes + # every XCom before an attempt, so a sub-task skipped on retry has them pushed again from here, + # next to its result, as a mapped task instance that succeeded keeps its own. + xcoms: dict[str, Any] | None = None + + def record_outlet_events(self, accessors: OutletEventAccessors) -> None: + """ + Keep a JSON-safe snapshot of the outlet asset events the indexed task recorded. + + Persisted with the checkpoint so a later attempt can :meth:`replay_outlet_events` when + the indexed task is skipped because it already succeeded. + """ + events: list[dict[str, Any]] = [] + for _asset_or_alias, accessor in accessors.items(): + if isinstance(accessor.key, AssetUniqueKey): + events.append( + { + "kind": "asset", + "name": accessor.key.name, + "uri": accessor.key.uri, + "extra": accessor.extra, + "partition_keys": sorted(accessor.partition_keys), + } + ) + for alias_event in accessor.asset_alias_events: + events.append( + { + "kind": "asset_alias", + "source_alias_name": alias_event.source_alias_name, + "dest_asset_key": { + "name": alias_event.dest_asset_key.name, + "uri": alias_event.dest_asset_key.uri, + }, + "dest_asset_extra": alias_event.dest_asset_extra, + "extra": alias_event.extra, + } + ) + if events: + self.outlet_events = events + + def replay_outlet_events(self, target: OutletEventAccessorsProtocol) -> None: + """ + Re-emit into ``target`` the events the indexed task recorded on a previous attempt. + + An indexed task skipped on retry (because it already succeeded) never re-executes, so it + never re-emits into the fresh ``OutletEventAccessors`` created for the new attempt. The + failed attempt sent nothing to the server either (outlet events travel only on the success + payload), so replaying cannot emit an event twice; without it the events would be lost. + They land in the parent's accessors as :meth:`IndexedTaskRunner.merge_outlet_events_into` + lands the live ones: one event per asset per task instance. + """ + for event in self.outlet_events or (): + if event["kind"] == "asset": + accessor = target[Asset(name=event["name"], uri=event["uri"])] + accessor.extra.update(event["extra"]) + if event["partition_keys"]: + accessor.add_partitions(event["partition_keys"]) + else: + target[AssetAlias(name=event["source_alias_name"])].asset_alias_events.append( + AssetAliasEvent( + source_alias_name=event["source_alias_name"], + dest_asset_key=AssetUniqueKey(**event["dest_asset_key"]), + dest_asset_extra=event["dest_asset_extra"], + extra=event["extra"], + ) + ) + + @staticmethod + def build_key(index: int) -> str: + # The task state store is already scoped to the parent task instance (dag, run, task and + # map index), so the key carries no identity, only a namespace that keeps the operator's own + # entries apart from anything user code stores from inside a sub-task. + return f"_iterable_{index}" + + def serialize(self) -> dict[str, Any]: + # The checkpoint travels to the supervisor as a JsonValue, which rejects anything that is not + # plain JSON (tuples, datetimes, models, ...). Serde turns those into JSON-compatible + # structures and restores them on read, exactly as XCom does with the same result. + data: dict[str, Any] = {"status": self.status.value} + if self.result is not None: + data["result"] = serde_serialize(self.result) + if self.outlet_events: + data["outlet_events"] = self.outlet_events + if self.fingerprint: + data["fingerprint"] = self.fingerprint + if self.try_number: + data["try_number"] = self.try_number + if self.xcoms: + data["xcoms"] = {key: serde_serialize(value) for key, value in self.xcoms.items()} + return data + + @classmethod + def deserialize(cls, raw: Any) -> IndexedTaskState | None: + if not isinstance(raw, Mapping): + return None + return cls( + status=TaskInstanceState(raw["status"]), + result=serde_deserialize(raw.get("result")), + outlet_events=raw.get("outlet_events"), + fingerprint=raw.get("fingerprint"), + try_number=raw.get("try_number", 0), + xcoms={key: serde_deserialize(value) for key, value in raw["xcoms"].items()} + if raw.get("xcoms") + else None, + ) + + +class IndexedTaskInstance(RuntimeTaskInstance): + """ + Indexed task instance to run a mapped operator. + + It shares the parent task instance's identity, so what an iteration pushes or stores lands in + the parent's scope, suffixed with the index so that iterations never overwrite each other: + XComs through :meth:`xcom_push`, task state through :attr:`task_state_store`. Reading back + follows the same rule: :meth:`xcom_pull` of the iteration's own XComs adds the index, a pull + from another task does not, as the store's accessor adds it to every key of its own. The + operator's own checkpoints go to the parent's store unsuffixed, under their + ``_iterable_`` keys. + """ + + index: int + parent_task_state_store: TaskStateStoreAccessor + input_fingerprint: str | None = None + # The XComs the sub-task pushed other than its return value, by their key before the index + # suffix: kept in memory while it runs and written once, with its SUCCESS checkpoint. + pushed_xcoms: dict[str, Any] = Field(default_factory=dict) + + @classmethod + def create_indexed_task( + cls, *, context: Context, index: int, operator: BaseOperator, input_fingerprint: str | None = None + ) -> IndexedTaskInstance: + """ + Create the runtime instance for one index of an iterated task from the parent's context. + + The instance shares the parent task instance's identity (id, run, map index, try number and + retry budget), so XComs and task state land in the parent's scope, and carries the unmapped + operator for that index. The budget is the parent's ``max_tries`` rather than the operator's + ``retries``: a manual clear raises it, and it is the parent Airflow retries. The parent's + state store comes from the context, the same accessor the operator's + checkpoints use. ``model_construct`` skips Pydantic validation on purpose: one instance is built per + indexed task, and the parent was validated already, so only the index needs checking here. + """ + if index < 0: + raise ValueError(f"IndexedTaskInstance requires index >= 0, got {index}") + parent = context["ti"] + return cls.model_construct( + id=parent.id, + parent_task_state_store=context["task_state_store"], + # The parent's: an iteration's XComs and state belong to the task instance that runs it. + task_id=parent.task_id, + dag_id=operator.dag_id, + run_id=parent.run_id, + map_index=parent.map_index, + index=index, + input_fingerprint=input_fingerprint, + max_tries=parent.max_tries, + start_date=parent.start_date, + state=TaskInstanceState.SCHEDULED.value, + is_mapped=True, + task=operator, + try_number=parent.try_number, + # The dag run, logical date and the rest the server sent for the task instance that + # runs this iteration: its template context, get_previous_ti() and the like read them. + _ti_context_from_server=getattr(parent, "_ti_context_from_server", None), + ) + + def xcom_push( + self, + key: str, + value: Any, + ): + super().xcom_push(key=f"{key}_{self.index}", value=value) + self._record_push(key, value) + + async def axcom_push( + self, + key: str, + value: Any, + ): + await super().axcom_push(key=f"{key}_{self.index}", value=value) + self._record_push(key, value) + + def _record_push(self, key: str, value: Any) -> None: + # The return value has its own slot on the checkpoint and is published by the operator. + if key != BaseXCom.XCOM_RETURN_KEY: + self.pushed_xcoms[key] = value + + def _own_key(self, task_ids: str | Iterable[str] | None, dag_id: str | None, key: str) -> str: + """ + Suffix ``key`` with the index for a pull of this iteration's own XComs. + + What :meth:`xcom_push` wrote under ``_`` is read back under the same name when + the pull names no task or this task instance's own (in its own DAG); a pull from another + task, or from several, keeps its key, since those XComs carry no index. + """ + own_task = task_ids is None or task_ids == self.task_id + own_dag = dag_id is None or dag_id == self.dag_id + return f"{key}_{self.index}" if own_task and own_dag else key + + def xcom_pull( + self, + task_ids: str | Iterable[str] | None = None, + dag_id: str | None = None, + key: str = BaseXCom.XCOM_RETURN_KEY, + include_prior_dates: bool = False, + *, + map_indexes: int | Iterable[int] | None | ArgNotSet = NOTSET, + default: Any = None, + run_id: str | None = None, + ) -> Any: + return super().xcom_pull( + task_ids=task_ids, + dag_id=dag_id, + key=self._own_key(task_ids, dag_id, key), + include_prior_dates=include_prior_dates, + map_indexes=map_indexes, + default=default, + run_id=run_id, + ) + + async def axcom_pull( + self, + task_ids: str | Iterable[str] | None = None, + dag_id: str | None = None, + key: str = BaseXCom.XCOM_RETURN_KEY, + include_prior_dates: bool = False, + *, + map_indexes: int | Iterable[int] | None | ArgNotSet = NOTSET, + default: Any = None, + run_id: str | None = None, + ) -> Any: + return await super().axcom_pull( + task_ids=task_ids, + dag_id=dag_id, + key=self._own_key(task_ids, dag_id, key), + include_prior_dates=include_prior_dates, + map_indexes=map_indexes, + default=default, + run_id=run_id, + ) + + @cached_property + def task_state_store(self) -> TaskStateStoreAccessor: # type: ignore[override] + """The parent's store seen from this iteration: keys are suffixed with the index.""" + return IndexedTaskStateStoreAccessor(self.parent_task_state_store, self.index) + + def context_for(self, context: Context, *, outlet_events: OutletEventAccessors | None = None) -> Context: + """ + Return the parent's context as this indexed task sees it. + + A clone of ``context`` with this task instance under ``ti`` and ``task_instance``, its + unmapped operator under ``task`` and its indexed view of the store under + ``task_state_store``: the keys ``context_update_for_unmapped`` sets for a mapped task + instance, and the store next to them, so that a template and ``execute`` read the same + store. The one place these keys are listed: :meth:`IndexedTaskRunner.indexed_context` + builds the context the task runs in from it, and ``IterableOperator._create_task`` the + one its templates are rendered against. ``outlet_events`` is swapped when given, for the + run; before it, nothing is emitted and the parent's accessor stays. + """ + indexed: Context = { + **clone_context(context), + "ti": self, + "task_instance": self, + "task": self.task, + "task_state_store": self.task_state_store, + } + if outlet_events is not None: + indexed["outlet_events"] = outlet_events + return indexed + + async def aget_state(self) -> IndexedTaskState | None: + return IndexedTaskState.deserialize(await self.parent_task_state_store.aget(self.state_key)) + + async def aset_state(self, state: IndexedTaskState) -> None: + await self.parent_task_state_store.aset(self.state_key, state.serialize()) + + @property + def is_async(self) -> bool: + return self.task.is_async + + @property + def is_eligible_to_retry(self) -> bool: + """ + Whether Airflow runs the parent again after this attempt fails. + + The same rule the API server applies to the parent (``_is_eligible_to_retry``), so the + callbacks of an iteration agree with what happens to the task instance. + """ + return self.max_tries != 0 and self.try_number <= self.max_tries + + @property + def state_key(self) -> str: + return IndexedTaskState.build_key(self.index) + + @property + def do_xcom_push(self) -> bool: + return self.task.do_xcom_push + + +class SubOperatorRegister(Protocol): + """Where an indexed task's operator is noted while its code runs, so that a kill reaches it.""" + + def register(self, operator: BaseOperator) -> None: ... + + def unregister(self, operator: BaseOperator) -> None: ... + + +class IndexedTaskRunner(LoggingMixin): + """ + Run one indexed task of an iterated task: its operator, against its own view of the context. + + Named apart from Airflow's executors, which schedule task instances, and from the task runner + process, which runs the parent: this runs one index inside that process, sync or async. + """ + + def __init__( + self, + task_instance: IndexedTaskInstance, + register: SubOperatorRegister | None = None, + outlet_events: OutletEventAccessors | None = None, + ): + """ + Run an operator or trigger for one sub-task instance. + + :param outlet_events: The accessor the sub-task's asset events are collected in, its own + so they can be checkpointed and merged apart from its siblings'. Created here when not + given; the caller reads it back through :attr:`outlet_events` after the run. + :param register: Optional register the operator is entered in while its code runs (see + :meth:`in_flight`), so that IterableOperator.on_kill() can reach whichever sub-tasks + are in flight; the iterated task's ``IterationState``. + """ + super().__init__() + self.task_instance = task_instance + self.outlet_events = outlet_events if outlet_events is not None else OutletEventAccessors() + self._result: Any | None = None + self._start_time: float | None = None + self._context: Context | None = None + self._register = register + #: The exception this indexed task failed with, noted by __exit__ and reported by + #: :meth:`report_failure` once the whole task's fate is known. + self.failure: BaseException | None = None + #: Whether the ``AirflowTaskTimeout`` this indexed task ended with was raised by its operator + #: off the main thread, where the parent's limit never strikes; see :meth:`in_flight`. + self.timed_out_on_its_own = False + self._cancelled = False + + def merge_outlet_events_into(self, target: OutletEventAccessorsProtocol) -> None: + """ + Fold the outlet asset events this indexed task recorded into the parent's ``target``. + + Called right after the indexed task succeeds, on the IterableOperator's shared + ``context["outlet_events"]``. A task instance sends one event per asset, so indexed tasks + that emit to the same asset end up in that one event: ``extra`` keeps what the last one to + finish wrote, while partition keys and alias events accumulate. ``.expand()`` sends one + event per mapped task instance instead; the docs page says so in its comparison table. + """ + for asset_or_alias, accessor in self.outlet_events.items(): + target_accessor = target[asset_or_alias] + target_accessor.extra.update(accessor.extra) + target_accessor.asset_alias_events.extend(accessor.asset_alias_events) + target_accessor.partition_keys.update(accessor.partition_keys) + + def cancel(self) -> None: + """ + Note that the coroutine waiting for this sync indexed task was cancelled. + + Its thread may go on, but the task gets no checkpoint from here on, and no report either: + the next attempt runs it again and reports it then. Best effort: a thread already past the + check in its exit still notes a failure. + """ + self._cancelled = True + + @property + def cancelled(self) -> bool: + """Whether :meth:`cancel` was called: the indexed task's exit then notes nothing.""" + return self._cancelled + + @property + def dag_id(self) -> str: + return self.task_instance.dag_id + + @property + def task_id(self) -> str: + return self.task_instance.task_id + + @property + def task_index(self) -> int: + return self.task_instance.index + + @property + def operator(self) -> BaseOperator: + return self.task_instance.task + + @property + def is_async(self) -> bool: + return self.task_instance.is_async + + @contextmanager + def indexed_context(self, context: Context) -> Iterator[Context]: + """ + Enter the parent's context as this indexed task sees it. + + Yields the task instance's :meth:`IndexedTaskInstance.context_for` view of the parent's + context with this runner's own outlet events, remembered on the runner and made the + current context for the duration of the block, so user code reads the indexed task's + unmapped operator under ``context["task"]``, not the IterableOperator. The parent's + context is left untouched: ``context_update_for_unmapped`` sets ``ti.task`` on whatever + ``ti`` it finds, which must be this task's, not the parent's. + """ + indexed_context = self.task_instance.context_for(context, outlet_events=self.outlet_events) + self._context = indexed_context + with set_indexed_context(indexed_context): + yield indexed_context + + @contextmanager + def in_flight(self) -> Iterator[None]: + """ + Register the operator as running for exactly as long as its code runs. + + Entered by :meth:`run` and :meth:`arun`, so in the worker thread or coroutine that executes + the operator: a sync operator stays registered while its thread is still inside ``execute``, + even after the coroutine waiting for it was cancelled, and ``IterableOperator.on_kill`` can + reach it. The operator the parent's execution timeout strikes stays registered as well: + the timeout lands on the main thread, where the loop runs async operators, and its + ``on_kill`` must not run there, where a synchronous SDK call raises, so + ``IterableOperator._run_tasks`` kills it off the loop thread with the others once the + timeout has unwound. ``TimeoutPosix`` raises on the main thread only, so an + ``AirflowTaskTimeout`` on a worker thread was raised by the operator itself (a hook that + gave up waiting): that is this indexed task's own failure, noted as + :attr:`timed_out_on_its_own`, and the operator is unregistered like any other. + """ + if self._register is not None: + self._register.register(self.operator) + struck_by_the_parent = False + try: + yield + except AirflowTaskTimeout: + self.timed_out_on_its_own = threading.current_thread() is not threading.main_thread() + struck_by_the_parent = not self.timed_out_on_its_own + raise + finally: + if self._register is not None and not struck_by_the_parent: + self._register.unregister(self.operator) + + def run(self, context: Context): + """Run the operator synchronously against this indexed task's own view of ``context``.""" + with self.in_flight(), self.indexed_context(context) as indexed_context: + return _execute_task(indexed_context, self.task_instance, self.log) + + async def arun(self, context: Context): + """Run the async operator against this indexed task's own view of ``context``.""" + with self.in_flight(), self.indexed_context(context) as indexed_context: + return await _execute_async_task(indexed_context, self.task_instance, self.log) + + def __enter__(self): + self._start_time = time.monotonic() + + if self.log.isEnabledFor(logging.INFO): + self.log.info( + "Running attempt %s of %s for %s with index %s in %s mode.", + self.task_instance.try_number, + self.task_instance.max_tries + 1, + self.task_instance.task_id, + self.task_index, + "async" if self.is_async else "sync", + ) + return self + + def __exit__(self, exc_type, exc_value, traceback): + elapsed = time.monotonic() - self._start_time if self._start_time else 0.0 + + if self._cancelled: + return None + if exc_value: + # Cancelled because the task is stopping, for a reason another iteration raised: this + # iteration neither failed nor will be retried on its own account, so it gets no state + # and no callback. The iteration that stopped the task reports its own outcome. + if isinstance(exc_value, CancelledError): + raise exc_value + # A skip is reported by report_skip() once its checkpoint is written, like a success. + if isinstance(exc_value, AirflowSkipException): + raise exc_value + # A failure is only noted here. Whether it is retried is the whole task's fate, which + # the other iterations decide too (a sibling's AirflowFailException fails it without a + # retry), so IterableOperator reports it through report_failure() once all have run. + self.log.error( + "Task instance %s for %s failed on attempt %s in %.2f seconds due to: %s", + self.task_index, + self.task_instance.task_id, + self.task_instance.try_number, + elapsed, + exc_value, + ) + self.failure = exc_value + raise exc_value + + # The state and the success callback follow once the checkpoint is written: report_success(). + if self.log.isEnabledFor(logging.INFO): + self.log.info( + "Task instance %s for %s finished successfully on attempt %s in %.2f seconds", + self.task_index, + self.task_instance.task_id, + self.task_instance.try_number, + elapsed, + ) + + def report_success(self) -> None: + """ + Report this indexed task as succeeded: its state, ``end_date`` and ``on_success_callback``. + + Called by ``IterableOperator._run_task`` once the SUCCESS checkpoint is written, where the + indexed task's code ran. The callback then speaks for work a retry will not run again: a + checkpoint write that fails fires nothing, and the attempt that runs the indexed task again + reports it then; a result push that fails after the checkpoint is replayed from it, so the + callback fires once. + """ + self.task_instance.end_date = datetime.now(tz=UTC) + self.task_instance.state = TaskInstanceState.SUCCESS + if self._context is not None: + _run_task_state_change_callbacks( + self.task_instance.task, "on_success_callback", self._context, self.log + ) + + def report_skip(self) -> None: + """Report this indexed task as skipped: its state, ``end_date`` and ``on_skipped_callback``, once its SKIPPED checkpoint is written.""" + self.task_instance.end_date = datetime.now(tz=UTC) + self.task_instance.state = TaskInstanceState.SKIPPED + if self._context is not None: + _run_task_state_change_callbacks( + self.task_instance.task, "on_skipped_callback", self._context, self.log + ) + + def report_failure(self, task_will_retry: bool) -> None: + """ + Report this indexed task's failure as what happens to the whole task. + + Called once every iteration has run, so the state and the callback agree with the task: + ``UP_FOR_RETRY`` and ``on_retry_callback`` when it is retried, ``FAILED`` and + ``on_failure_callback`` otherwise. + """ + if task_will_retry: + self.task_instance.end_date = datetime.now(tz=UTC) + self.task_instance.state = TaskInstanceState.UP_FOR_RETRY + else: + self.task_instance.state = TaskInstanceState.FAILED + if self._context is not None: + _run_task_state_change_callbacks( + self.task_instance.task, + "on_retry_callback" if task_will_retry else "on_failure_callback", + self._context, + self.log, + ) + + def _xcom_push( ti: RuntimeTaskInstance, key: str, @@ -2189,6 +2789,8 @@ def _run_execute_callable( context: Context, execute: Callable[..., Any] | functools.partial[Any], task: BaseOperator, + *, + enforce_timeout: bool = True, ) -> Any: """ Run the task's execute callable, applying the execution timeout if one is set. @@ -2198,10 +2800,15 @@ def _run_execute_callable( than under the caller. ``ExecutorSafeguard``'s tracker is set into that copy so the operator's ``execute`` passes the safeguard check, while the copy keeps the change from leaking into the surrounding context. + + ``enforce_timeout`` is False for the indexed tasks of an iterated task: they run in worker + threads, where ``TimeoutPosix`` cannot fire, under the parent's own limit, which the + parent already told the supervisor about; sending ``SetExecutionTimeout`` again per indexed + task would move the supervisor's deadline to the last one started. """ ctx = contextvars.copy_context() ctx.run(ExecutorSafeguard.tracker.set, task) - if task.execution_timeout: + if task.execution_timeout and enforce_timeout: from airflow.sdk.execution_time.timeout import timeout # TODO: handle timeout in case of deferral @@ -2245,10 +2852,45 @@ def _execute_task(context: Context, ti: RuntimeTaskInstance, log: Logger): assert isinstance(kwargs, dict) execute = functools.partial(task.resume_execution, next_method=next_method, next_kwargs=kwargs) - # Export context in os.environ to make it available for operators to use. - airflow_context_vars = context_to_airflow_vars(context, in_env_var_format=True) - os.environ.update(airflow_context_vars) + # Export the context to os.environ for operators that read AIRFLOW_CTX_* directly. Indexed + # sub-tasks skip this: they run concurrently in one process, so the update would race, and the + # parent IterableOperator already exported the same values before they started. + if not isinstance(ti, IndexedTaskInstance): + os.environ.update(context_to_airflow_vars(context, in_env_var_format=True)) + + outlet_events = _run_pre_execute(task, context, log) + log.info("::endgroup::") + + # An indexed sub-task runs in a worker thread under the parent's execution_timeout, which the + # parent enforces and reported to the supervisor once (see _run_execute_callable). + result = _run_execute_callable( + context, execute, task, enforce_timeout=not isinstance(ti, IndexedTaskInstance) + ) + + _run_post_execute(task, context, outlet_events, result, log) + return result + + +async def _execute_async_task(context: Context, ti: RuntimeTaskInstance, log: Logger): + """Async counterpart of :func:`_execute_task` for :class:`BaseAsyncOperator` sub-tasks.""" + task = cast("BaseAsyncOperator", ti.task) + # Async tasks cannot be resuming a deferral, so next_method never applies here. + outlet_events = _run_pre_execute(task, context, log) + + ctx = contextvars.copy_context() + ctx.run(ExecutorSafeguard.tracker.set, task) + # Under the parent's execution_timeout only, as a sync indexed task is: the operator's own limit + # is the same value, started later, so enforcing it here too would never fire first and would + # race the parent's kill when the two coincide. + result = await ctx.run(lambda: task.aexecute(context=context)) + + _run_post_execute(task, context, outlet_events, result, log) + return result + + +def _run_pre_execute(task: BaseOperator, context: Context, log: Logger) -> OutletEventAccessorsProtocol: + """Run the pre-execute hooks and the on_execute callback; shared by the sync and async paths.""" outlet_events = context_get_outlet_events(context) if (pre_execute_hook := task._pre_execute_hook) is not None: @@ -2257,18 +2899,22 @@ def _execute_task(context: Context, ti: RuntimeTaskInstance, log: Logger): create_executable_runner(pre_execute_hook, outlet_events, logger=log).run(context) _run_task_state_change_callbacks(task, "on_execute_callback", context, log) + return outlet_events - log.info("::endgroup::") - - result = _run_execute_callable(context, execute, task) +def _run_post_execute( + task: BaseOperator, + context: Context, + outlet_events: OutletEventAccessorsProtocol, + result: Any, + log: Logger, +) -> None: + """Run the post-execute hooks; shared by the sync and async paths.""" if (post_execute_hook := task._post_execute_hook) is not None: create_executable_runner(post_execute_hook, outlet_events, logger=log).run(context, result) if getattr(post_execute_hook := task.post_execute, "__func__", None) is not BaseOperator.post_execute: create_executable_runner(post_execute_hook, outlet_events, logger=log).run(context) - return result - def _render_map_index(context: Context, ti: RuntimeTaskInstance, log: Logger) -> str | None: """Render named map index if the Dag author defined map_index_template at the task level.""" diff --git a/task-sdk/tests/task_sdk/bases/test_decorator.py b/task-sdk/tests/task_sdk/bases/test_decorator.py index b3dfd1630c5d4..ab8ec74efb618 100644 --- a/task-sdk/tests/task_sdk/bases/test_decorator.py +++ b/task-sdk/tests/task_sdk/bases/test_decorator.py @@ -24,7 +24,7 @@ import pytest -from airflow.sdk import task +from airflow.sdk import DAG, task from airflow.sdk.bases.decorator import KNOWN_CONTEXT_KEYS, DecoratedOperator, is_async_callable RAW_CODE = """ @@ -413,3 +413,112 @@ def sync_task_fn(): return 42 assert not is_async_callable(sync_task_fn) + + +class TestTaskDecoratorTaskConcurrency: + """task_concurrency only has meaning for Dynamic Task Iteration (as the sub-task thread count + consumed by IterableOperator via .iterate()/.iterate_kwargs()). A plain + .expand()/.expand_kwargs() on a @task-decorated function never reaches that code path, so it + must be rejected instead of silently accepted as a dead value -- mirroring OperatorPartial.""" + + def test_direct_call_rejects_task_concurrency(self): + """Calling a @task-decorated function directly (no .expand()/.iterate()) constructs the + operator right away via BaseOperator.__init__, so task_concurrency must be rejected.""" + with DAG("test_dag"): + + @task(task_concurrency=2) + def add_one(x): + return x + 1 + + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + add_one(1) + + def test_expand_rejects_task_concurrency(self): + with DAG("test_dag"): + + @task(task_concurrency=2) + def add_one(x): + return x + 1 + + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + add_one.expand(x=[1, 2, 3]) + + def test_expand_kwargs_rejects_task_concurrency(self): + with DAG("test_dag"): + + @task(task_concurrency=2) + def add_one(x): + return x + 1 + + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + add_one.expand_kwargs([{"x": 1}, {"x": 2}]) + + def test_iterate_accepts_task_concurrency(self): + """.iterate() is the one entry point where task_concurrency is meaningful: it produces an + IterableOperator, which reads task_concurrency out of partial_kwargs as max_workers rather + than forwarding it to BaseOperator.__init__.""" + from airflow.sdk.definitions.iterableoperator import IterableOperator + + with DAG("test_dag"): + + @task(task_concurrency=2) + def add_one(x): + return x + 1 + + xcom_arg = add_one.iterate(x=[1, 2, 3]) + + assert isinstance(xcom_arg.operator, IterableOperator) + assert xcom_arg.operator.max_workers == 2 + + def test_iterate_kwargs_accepts_task_concurrency(self): + """.iterate_kwargs() is the list-of-dicts counterpart to .iterate() and must accept + task_concurrency the same way.""" + from airflow.sdk.definitions.iterableoperator import IterableOperator + + with DAG("test_dag"): + + @task(task_concurrency=2) + def add_one(x): + return x + 1 + + xcom_arg = add_one.iterate_kwargs([{"x": 1}, {"x": 2}]) + + assert isinstance(xcom_arg.operator, IterableOperator) + assert xcom_arg.operator.max_workers == 2 + + +def test_iterate_without_arguments_names_iterate_in_the_error(): + with DAG("test_dag"): + + @task + def add_one(x): + return x + 1 + + with pytest.raises(TypeError, match="no arguments to iterate against"): + add_one.iterate() + + +def test_iterate_ignores_multiple_outputs_inferred_from_return_annotation(): + with DAG("test_dag"): + + @task + def to_dict(x) -> dict: + return {"x": x} + + assert to_dict.multiple_outputs is True + xcom_arg = to_dict.iterate(x=[1, 2, 3]) + + assert xcom_arg.operator.multiple_outputs is False + + +def test_task_protocol_declares_every_method_of_a_decorated_callable(): + """ + ``@task`` is typed as returning ``Task``, so a public method ``_TaskDecorator`` has and ``Task`` + does not declare, such as ``.iterate()``, is an attribute error under a type checker. + """ + from airflow.sdk.bases.decorator import Task, _TaskDecorator + + def public_methods(cls): + return {name for name, value in vars(cls).items() if not name.startswith("_") and callable(value)} + + assert public_methods(_TaskDecorator) <= public_methods(Task) diff --git a/task-sdk/tests/task_sdk/bases/test_operator.py b/task-sdk/tests/task_sdk/bases/test_operator.py index d2060362f2b94..6e49caa136c06 100644 --- a/task-sdk/tests/task_sdk/bases/test_operator.py +++ b/task-sdk/tests/task_sdk/bases/test_operator.py @@ -20,6 +20,7 @@ import asyncio import copy import logging +import os import uuid import warnings from datetime import UTC, date, datetime, timedelta @@ -1291,3 +1292,69 @@ async def test_refuses_running_loop(self): pass assert not running.is_closed() + + +class TestTaskConcurrency: + """ + ``task_concurrency`` is read by ``.iterate()`` only and rejected when passed anywhere else. + + ``default_args`` never set it: Airflow 2 used the name for what is now ``max_active_tis_per_dag``, + and a DAG that still carries it in its ``default_args`` parsed on Airflow 3 before iterated tasks + existed, so it must keep parsing. + """ + + @pytest.fixture + def airflow2_default_args(self): + return {"task_concurrency": 1} + + def test_direct_instantiation_rejects_it(self): + with DAG("test_dag"): + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + MockOperator(task_id="op", task_concurrency=2) + + def test_expand_rejects_it(self): + with DAG("test_dag"): + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + MockOperator.partial(task_id="op", task_concurrency=2).expand(arg1=["a", "b"]) + + def test_iterate_reads_it(self): + with DAG("test_dag"): + iterated = MockOperator.partial(task_id="op", task_concurrency=2).iterate(arg1=["a", "b"]) + + assert iterated.max_workers == 2 + + def test_dag_default_args_do_not_reject_an_operator(self, airflow2_default_args): + with DAG("test_dag", default_args=airflow2_default_args): + op = MockOperator(task_id="op") + + assert "task_concurrency" not in op._BaseOperator__init_kwargs + + def test_task_default_args_do_not_reject_an_operator(self, airflow2_default_args): + with DAG("test_dag"): + MockOperator(task_id="op", default_args=airflow2_default_args) + + def test_dag_default_args_do_not_reject_a_task_call(self, airflow2_default_args): + with DAG("test_dag", default_args=airflow2_default_args): + + @task_decorator + def add_one(x): + return x + 1 + + add_one(1) + + def test_dag_default_args_do_not_reject_expand(self, airflow2_default_args): + with DAG("test_dag", default_args=airflow2_default_args): + mapped = MockOperator.partial(task_id="op").expand(arg1=["a", "b"]) + + assert "task_concurrency" not in mapped.partial_kwargs + + def test_dag_default_args_do_not_set_it_for_iterate(self, airflow2_default_args): + with DAG("test_dag", default_args=airflow2_default_args): + iterated = MockOperator.partial(task_id="op").iterate(arg1=["a", "b"]) + + assert iterated.max_workers == (os.cpu_count() or 1) + + def test_iterate_does_not_take_it_as_an_iterated_argument(self): + with DAG("test_dag"): + with pytest.raises(TypeError, match="unexpected keyword argument 'task_concurrency'"): + MockOperator.partial(task_id="op").iterate(arg1=["a"], task_concurrency=[1, 2]) diff --git a/task-sdk/tests/task_sdk/bases/test_xcom.py b/task-sdk/tests/task_sdk/bases/test_xcom.py index caff4df089afc..71433a49dc2d3 100644 --- a/task-sdk/tests/task_sdk/bases/test_xcom.py +++ b/task-sdk/tests/task_sdk/bases/test_xcom.py @@ -18,10 +18,11 @@ from __future__ import annotations from unittest import mock +from unittest.mock import AsyncMock, patch import pytest -from airflow.sdk.bases.xcom import BaseXCom +from airflow.sdk.bases.xcom import BaseXCom, XComIterable from airflow.sdk.execution_time.comms import ( DeleteXCom, GetXCom, @@ -29,6 +30,7 @@ XComResult, XComSequenceSliceResult, ) +from airflow.sdk.execution_time.xcom import XCom from airflow.sdk.types import TaskInstanceKey @@ -288,3 +290,176 @@ async def test_aget_value_calls_aget_one(self, mock_supervisor_comms): include_prior_dates=False, ) ) + + +class TestXComIterable: + def make_iterable(self, length: int = 0, map_index: int | None = None) -> XComIterable: + return XComIterable(task_id="task", dag_id="dag", run_id="run", map_index=map_index, length=length) + + def test_has_no_append(self): + """The consumer-facing Sequence is read-only: nothing on it mutates the underlying XComs.""" + iterable = self.make_iterable(length=1) + assert not hasattr(iterable, "append") + assert not hasattr(iterable, "aappend") + + def test_serialize_returns_expected_dict(self): + iterable = self.make_iterable(length=3, map_index=1) + assert iterable.serialize() == { + "task_id": "task", + "dag_id": "dag", + "run_id": "run", + "map_index": 1, + "length": 3, + "skipped": [], + } + + def test_deserialize_restores_fields(self): + data = {"task_id": "task", "dag_id": "dag", "run_id": "run", "map_index": 2, "length": 5} + iterable = XComIterable.deserialize(data, version=1) + assert iterable.task_id == "task" + assert iterable.dag_id == "dag" + assert iterable.run_id == "run" + assert iterable.map_index == 2 + assert iterable.length == 5 + assert iterable.skipped == [] + assert len(iterable) == 5 + + def test_skipped_indices_round_trip(self): + iterable = XComIterable(task_id="task", dag_id="dag", run_id="run", length=4, skipped=[3, 1]) + restored = XComIterable.deserialize(iterable.serialize(), version=1) + assert restored.skipped == [1, 3] + assert len(restored) == 2 + + @pytest.mark.asyncio + @patch.object(XCom, "aget_one", new_callable=AsyncMock, return_value="value-1") + async def test_aget_calls_xcom_aget_one_with_indexed_key(self, mock_aget_one): + iterable = self.make_iterable(length=2, map_index=3) + assert await iterable.aget(1) == "value-1" + mock_aget_one.assert_awaited_once_with( + key=f"{BaseXCom.XCOM_RETURN_KEY}_1", + dag_id="dag", + task_id="task", + run_id="run", + map_index=3, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("index", [-3, 2]) + @patch.object(XCom, "aget_one", new_callable=AsyncMock) + async def test_aget_out_of_range_raises_index_error_without_fetching(self, mock_aget_one, index): + iterable = self.make_iterable(length=2) + with pytest.raises(IndexError): + await iterable.aget(index) + mock_aget_one.assert_not_awaited() + + @pytest.mark.asyncio + @patch.object(XCom, "aget_one", new_callable=AsyncMock, return_value="last") + async def test_aget_negative_index_counts_from_the_end(self, mock_aget_one): + iterable = self.make_iterable(length=3) + assert await iterable.aget(-1) == "last" + assert mock_aget_one.await_args.kwargs["key"] == f"{BaseXCom.XCOM_RETURN_KEY}_2" + + @pytest.mark.asyncio + @patch.object(XCom, "get_one") + @patch.object(XCom, "aget_one", new_callable=AsyncMock) + async def test_async_iteration_reads_every_item_in_order_through_aget_one( + self, mock_aget_one, mock_get_one + ): + """``async for`` never touches the synchronous ``get_one``, so it is safe on the task's event loop.""" + mock_aget_one.side_effect = self._pages_by_key(["a", "b", "c"]) + iterable = self.make_iterable(length=3) + assert [item async for item in iterable] == ["a", "b", "c"] + assert [call.kwargs["key"] for call in mock_aget_one.await_args_list] == [ + f"{BaseXCom.XCOM_RETURN_KEY}_{index}" for index in range(3) + ] + mock_get_one.assert_not_called() + + @pytest.mark.asyncio + @patch.object(XCom, "aget_one", new_callable=AsyncMock) + async def test_async_iteration_on_empty_iterable_yields_nothing(self, mock_aget_one): + iterable = self.make_iterable(length=0) + assert [item async for item in iterable] == [] + mock_aget_one.assert_not_awaited() + + @staticmethod + def _pages_by_key(pages: list) -> object: + """ + Build a get_one side_effect that maps each page's index-suffixed key to its page. + + Unlike a plain list side_effect (consumed once and then exhausted), this can be called any + number of times for the same key, mirroring how a real XCom backend is queried by key and + does not get "used up". + """ + + def _get_one(*args, **kwargs): + index = int(kwargs["key"].rsplit("_", 1)[-1]) + page = pages[index] + return page() if callable(page) else page + + return _get_one + + @patch.object(XCom, "get_one") + def test_negative_indices_count_from_the_end(self, mock_get_one): + """The Sequence contract holds: ``[-1]`` is the last element.""" + mock_get_one.side_effect = self._pages_by_key([["a", "b"], ["c"]]) + iterable = self.make_iterable(length=2) + + assert iterable[-1] == ["c"] + assert iterable[-2] == ["a", "b"] + + @patch.object(XCom, "get_one") + def test_negative_indices_out_of_range_raise(self, mock_get_one): + mock_get_one.side_effect = self._pages_by_key([["a", "b"], ["c"]]) + iterable = self.make_iterable(length=2) + + with pytest.raises(IndexError): + iterable[-3] + + @staticmethod + def _values_by_index(values: dict[int, str]) -> object: + """A get_one/aget_one side_effect that returns the value pushed under each index-suffixed key.""" + + def _get_one(*args, **kwargs): + return values[int(kwargs["key"].rsplit("_", 1)[-1])] + + return _get_one + + def make_skipping_iterable(self) -> XComIterable: + """Five input items, of which indices 0, 2 and 3 were skipped: only 1 and 4 hold a value.""" + return XComIterable(task_id="task", dag_id="dag", run_id="run", length=5, skipped=[0, 2, 3]) + + @patch.object(XCom, "get_one") + def test_skipped_indices_are_left_out(self, mock_get_one): + mock_get_one.side_effect = self._values_by_index({1: "one", 4: "four"}) + iterable = self.make_skipping_iterable() + + assert len(iterable) == 2 + assert list(iterable) == ["one", "four"] + assert iterable[0] == "one" + assert iterable[1] == "four" + assert iterable[-1] == "four" + assert iterable[-2] == "one" + assert iterable[::-1] == ["four", "one"] + with pytest.raises(IndexError): + iterable[2] + with pytest.raises(IndexError): + iterable[-3] + assert sorted({call.kwargs["key"] for call in mock_get_one.call_args_list}) == [ + f"{BaseXCom.XCOM_RETURN_KEY}_1", + f"{BaseXCom.XCOM_RETURN_KEY}_4", + ] + + @pytest.mark.asyncio + @patch.object(XCom, "aget_one", new_callable=AsyncMock) + async def test_skipped_indices_are_left_out_when_read_asynchronously(self, mock_aget_one): + mock_aget_one.side_effect = self._values_by_index({1: "one", 4: "four"}) + iterable = self.make_skipping_iterable() + + assert await iterable.alen() == 2 + assert [item async for item in iterable] == ["one", "four"] + assert await iterable.aget(-1) == "four" + + def test_every_index_skipped_is_empty(self): + iterable = XComIterable(task_id="task", dag_id="dag", run_id="run", length=2, skipped=[0, 1]) + assert len(iterable) == 0 + assert list(iterable) == [] diff --git a/task-sdk/tests/task_sdk/definitions/_internal/test_expandinput.py b/task-sdk/tests/task_sdk/definitions/_internal/test_expandinput.py new file mode 100644 index 0000000000000..c988474521fea --- /dev/null +++ b/task-sdk/tests/task_sdk/definitions/_internal/test_expandinput.py @@ -0,0 +1,352 @@ +# 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 + +import asyncio +import threading +from collections.abc import Sequence + +import pytest +from task_sdk.definitions.conftest import make_xcom_arg + +from airflow.sdk.definitions._internal.expandinput import ( + DecoratedExpandInput, + DictOfListsExpandInput, + ExpandInput, + ListOfDictsExpandInput, + MappedArgument, + Resolved, + Source, + index_for_each_field, +) +from airflow.sdk.exceptions import UnmappableXComTypePushed, XComForMappingNotPushed + + +class AsyncOnlyValues(Sequence): + """A Sequence read through its own async accessors, like an XComIterable on the loop; sync reads fail.""" + + def __init__(self, values): + self.values = values + self.reads: list[int] = [] + + def __len__(self): + pytest.fail("synchronous __len__ used on the async path") + + def __getitem__(self, index): + pytest.fail("synchronous __getitem__ used on the async path") + + async def alen(self): + return len(self.values) + + async def aget(self, index): + self.reads.append(index) + return self.values[index] + + +def _async_only_xcom_arg(values): + """An XComArg whose sync ``resolve`` must never run and whose ``aresolve`` gives an async-read value.""" + xcom_arg = make_xcom_arg(None) + xcom_arg.resolve = lambda *a, **kw: pytest.fail("synchronous resolve() used on the async path") + + async def aresolve(*a, **kw): + return AsyncOnlyValues(values) + + xcom_arg.aresolve = aresolve + return xcom_arg + + +async def _items(expand_input: ExpandInput, context=None) -> list: + """Every item of the input in index order, the way IterableOperator.execute reads it.""" + length, aget = await expand_input.aresolve(context or {}) + return [await aget(index) for index in range(length)] + + +class TestExpandInput: + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("actual", "expected"), + [ + ({"a": [1, 2, 3]}, [{"a": 1}, {"a": 2}, {"a": 3}]), + ( + {"a": [1, 2], "b": [10, 20]}, + [{"a": 1, "b": 10}, {"a": 1, "b": 20}, {"a": 2, "b": 10}, {"a": 2, "b": 20}], + ), + ({"a": [1, 2], "b": [10, 20], "c": ["x"]}, None), + ({"a": (x for x in [1, 2])}, [{"a": 1}, {"a": 2}]), + ({"a": {"x": 1, "y": 2}}, [{"a": ("x", 1)}, {"a": ("y", 2)}]), + ({"a": []}, []), + ({"a": [1, 2], "b": []}, []), + ({"a": AsyncOnlyValues([1, 2])}, [{"a": 1}, {"a": 2}]), + ], + ) + async def test_dict_of_lists_expand_input_aresolve(self, actual, expected): + """ + The cross product in the order ``.expand()`` uses (the last argument varies fastest). A dict + value expands to its (key, value) pairs, mirroring _expand_mapped_field's handling of dict + values, so .iterate() and .expand() hand sub-tasks the same per-index value for a dict argument. + """ + expand_input = DictOfListsExpandInput(actual) + resolved = await expand_input.aresolve({}) + assert isinstance(resolved, Resolved) + + # Read from this resolution: a generator argument is consumed by the pull and cannot be resolved twice. + items = [await resolved.aget(index) for index in range(resolved.length)] + if expected is None: + expected = [ + {"a": a, "b": b, "c": "x"} for a in (1, 2) for b in (10, 20) + ] # 3 arguments: 2 x 2 x 1 combinations + assert items == expected + assert resolved.length == len(expected) + + @pytest.mark.asyncio + async def test_dict_of_lists_expand_input_aresolve_pulls_xcom_args_with_aresolve(self): + expand_input = DictOfListsExpandInput({"a": _async_only_xcom_arg([1, 2]), "b": [10, 20]}) + + assert await _items(expand_input) == [ + {"a": 1, "b": 10}, + {"a": 1, "b": 20}, + {"a": 2, "b": 10}, + {"a": 2, "b": 20}, + ] + + @pytest.mark.asyncio + async def test_dict_of_lists_expand_input_aresolve_reads_each_index_on_demand(self): + """Sources are pulled once up front, but items are read only as the consumer asks for them.""" + source = AsyncOnlyValues([0, 1, 2]) + expand_input = DictOfListsExpandInput({"a": source, "b": [10, 20]}) + length, aget = await expand_input.aresolve({}) + + assert length == 6 + assert source.reads == [] + assert await aget(0) == {"a": 0, "b": 10} + assert await aget(1) == {"a": 0, "b": 20} + assert source.reads == [0, 0] + assert await aget(2) == {"a": 1, "b": 10} + assert source.reads == [0, 0, 1] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("actual", "expected"), + [ + ([{"a": 1}, {"a": 2}], [{"a": 1}, {"a": 2}]), + ([{"a": 1, "b": 2}], [{"a": 1, "b": 2}]), + ([], []), + ], + ) + async def test_list_of_dicts_expand_input_aresolve(self, actual, expected): + expand_input = ListOfDictsExpandInput(actual) + assert await _items(expand_input) == expected + + @pytest.mark.asyncio + async def test_list_of_dicts_expand_input_aresolve_pulls_xcom_args_with_aresolve(self): + """An XComArg is one mapping per item, as for ``expand_kwargs``: the whole list, or one entry of it.""" + whole = ListOfDictsExpandInput(_async_only_xcom_arg([{"a": 1}, {"a": 2}])) + assert await _items(whole) == [{"a": 1}, {"a": 2}] + + entry = make_xcom_arg(None) + + async def aresolve(*a, **kw): + return {"a": 1} + + entry.aresolve = aresolve + per_item = ListOfDictsExpandInput([{"a": 0}, entry]) + assert await _items(per_item) == [{"a": 0}, {"a": 1}] + + @pytest.mark.asyncio + async def test_list_of_dicts_expand_input_aresolve_rejects_non_mapping_items(self): + expand_input = ListOfDictsExpandInput(_async_only_xcom_arg([{"a": 1}, 2])) + _, aget = await expand_input.aresolve({}) + + assert await aget(0) == {"a": 1} + with pytest.raises(ValueError, match=r"iterate_kwargs\(\) expects a list\[dict\], not list\[int\]"): + await aget(1) + + def test_decorated_expand_inputs_compare_by_their_delegate(self): + one = DecoratedExpandInput(ListOfDictsExpandInput([{"a": 1}])) + same = DecoratedExpandInput(ListOfDictsExpandInput([{"a": 1}])) + other = DecoratedExpandInput(ListOfDictsExpandInput([{"a": 2}])) + + assert one == same + assert one != other + + @pytest.mark.asyncio + async def test_decorated_expand_input_aresolve_wraps_op_kwargs(self): + decorated = DecoratedExpandInput(ListOfDictsExpandInput([{"a": 1}, {"a": 2}])) + + assert await _items(decorated) == [{"op_kwargs": {"a": 1}}, {"op_kwargs": {"a": 2}}] + + @pytest.mark.asyncio + async def test_base_aresolve_is_abstract(self): + class Incomplete(ExpandInput): + EXPAND_INPUT_TYPE = "incomplete" + + @property + def value(self): + return None + + with pytest.raises(NotImplementedError): + await Incomplete().aresolve({}) + + +class TestSource: + """One resolved expand argument, read by index the way ``_expand_mapped_field`` picks it.""" + + @pytest.mark.asyncio + async def test_mapping_is_read_as_its_items(self): + source = await Source.from_argument({"x": 1, "y": 2}, {}) + assert await source.alen() == 2 + assert [await source.aget(index) for index in range(2)] == [("x", 1), ("y", 2)] + + @pytest.mark.asyncio + @pytest.mark.parametrize("value", ["hello", b"bytes", 1, None, 2.5]) + async def test_scalars_and_strings_are_refused(self, value): + """A literal that is no collection is refused at parse time; one that gets here is refused too.""" + with pytest.raises(TypeError, match=f"cannot iterate over a '{type(value).__name__}' argument"): + await Source.from_argument(value, {}) + + @pytest.mark.asyncio + @pytest.mark.parametrize("value", [[1, 2], (1, 2), range(1, 3)]) + async def test_in_memory_sequences_are_read_in_place(self, value, monkeypatch): + """Nothing here can block on a supervisor call, so no read goes through a worker thread.""" + monkeypatch.setattr(asyncio, "to_thread", lambda *a, **kw: pytest.fail("to_thread used")) + source = await Source.from_argument(value, {}) + assert [await source.aget(index) for index in range(2)] == [1, 2] + + @pytest.mark.asyncio + async def test_other_iterables_are_materialized(self): + source = await Source.from_argument((x for x in [1, 2]), {}) + assert await source.alen() == 2 + assert await source.aget(1) == 2 + + @pytest.mark.asyncio + async def test_xcom_arg_is_pulled_with_aresolve_then_normalised(self): + source = await Source.from_argument(_async_only_xcom_arg([1, 2]), {}) + assert await source.alen() == 2 + assert await source.aget(1) == 2 + + xcom_arg = make_xcom_arg(None) + + async def aresolve(*a, **kw): + return {"x": 1} + + xcom_arg.aresolve = aresolve + source = await Source.from_argument(xcom_arg, {}) + assert await source.aget(0) == ("x", 1) + + @pytest.mark.asyncio + @pytest.mark.parametrize("value", ["hello", b"bytes", 1, 2.5, object()]) + async def test_an_unmappable_upstream_value_is_rejected(self, value): + """ + What ``.expand()`` refuses at push time is refused here at read time. + + An upstream with a mapped dependant raises ``UnmappableXComTypePushed`` from + ``_push_xcom_if_needed``, but an IterableOperator is no ``MappedOperator`` and is not found + by ``iter_mapped_dependants``, so that check never fires for it; without this one a string + from upstream would be iterated as a single item. + """ + with pytest.raises(UnmappableXComTypePushed, match=type(value).__name__): + await Source.from_argument(make_xcom_arg(value), {}) + + @pytest.mark.asyncio + async def test_an_upstream_that_pushed_nothing_is_rejected(self): + with pytest.raises(XComForMappingNotPushed): + await Source.from_argument(make_xcom_arg(None), {}) + + @pytest.mark.asyncio + async def test_async_accessors_are_preferred(self, monkeypatch): + monkeypatch.setattr(asyncio, "to_thread", lambda *a, **kw: pytest.fail("to_thread used")) + value = AsyncOnlyValues([1, 2]) + source = await Source.from_argument(value, {}) + assert await source.alen() == 2 + assert await source.aget(1) == 2 + assert value.reads == [1] + + @pytest.mark.asyncio + async def test_sequences_without_async_accessors_are_read_off_the_loop_thread(self): + """ + A sequence may block in ``__len__``/``__getitem__`` with a synchronous supervisor call (a + ``.map()`` result over a mapped upstream), which on the loop thread would deadlock with the + ``asend`` calls in flight, so it is read in a worker thread. + """ + threads: list[threading.Thread] = [] + + class Blocking: + def __len__(self): + threads.append(threading.current_thread()) + return 2 + + def __getitem__(self, index): + threads.append(threading.current_thread()) + return [1, 2][index] + + source = Source(Blocking()) + assert await source.alen() == 2 + assert await source.aget(1) == 2 + assert len(threads) == 2 + assert all(thread is not threading.current_thread() for thread in threads) + assert asyncio.get_running_loop().is_running() + + +class TestIndexForEachField: + """The one cross-product rule behind ``_expand_mapped_field`` (``.expand()``) and ``aresolve`` (``.iterate()``).""" + + @pytest.mark.parametrize( + ("map_index", "expected"), + [ + (0, {"a": 0, "b": 0, "c": 0}), + (1, {"a": 0, "b": 0, "c": 1}), + (2, {"a": 0, "b": 1, "c": 0}), + (5, {"a": 0, "b": 2, "c": 1}), + (6, {"a": 1, "b": 0, "c": 0}), + (11, {"a": 1, "b": 2, "c": 1}), + ], + ) + def test_last_argument_varies_fastest(self, map_index, expected): + assert index_for_each_field(map_index, {"a": 2, "b": 3, "c": 2}) == expected + + def test_single_argument_is_the_position_itself(self): + assert index_for_each_field(4, {"a": 9}) == {"a": 4} + + def test_matches_itertools_product_order(self): + import itertools + + lengths = {"a": 2, "b": 3, "c": 2} + positions = [tuple(index_for_each_field(i, lengths).values()) for i in range(12)] + assert positions == list(itertools.product(range(2), range(3), range(2))) + + def test_zero_length_argument_cannot_be_expanded(self): + with pytest.raises(RuntimeError, match="cannot expand field mapped to length 0"): + index_for_each_field(0, {"a": 2, "b": 0}) + + @pytest.mark.asyncio + async def test_expand_and_iterate_pick_the_same_item_for_a_position(self): + """``resolve`` at ``map_index`` and ``aresolve``'s ``aget`` at that index agree, dict argument included.""" + expand_input = DictOfListsExpandInput({"a": [1, 2], "b": {"x": 10, "y": 20, "z": 30}}) + _, aget = await expand_input.aresolve({}) + for map_index in range(6): + ti = type("TI", (), {"map_index": map_index, "_upstream_map_indexes": {}})() + expanded, _ = expand_input.resolve({"ti": ti}) + assert await aget(map_index) == dict(expanded) + + +def test_mapped_argument_is_keyword_only(): + """``MappedArgument`` takes its input and key by keyword, as on main.""" + expand_input = DictOfListsExpandInput({"a": [1, 2]}) + + with pytest.raises(TypeError): + MappedArgument(expand_input, "a") # type: ignore[misc] + assert MappedArgument(input=expand_input, key="a") == MappedArgument(input=expand_input, key="a") diff --git a/task-sdk/tests/task_sdk/definitions/conftest.py b/task-sdk/tests/task_sdk/definitions/conftest.py index 3f89f34b4d2da..c63b9cbe06be5 100644 --- a/task-sdk/tests/task_sdk/definitions/conftest.py +++ b/task-sdk/tests/task_sdk/definitions/conftest.py @@ -17,11 +17,12 @@ # under the License. from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any import pytest import structlog +from airflow.sdk import BaseOperator, XComArg from airflow.sdk.execution_time.comms import SucceedTask, TaskState if TYPE_CHECKING: @@ -37,6 +38,12 @@ def run(dag: DAG, task_id: str, map_index: int): log = structlog.get_logger(__name__) mock_supervisor_comms.send.reset_mock() + # Tests queue their supervisor replies on the sync ``send``. Answer the async ``asend`` from the + # same replies, so a task pulling through the async SDK path (iterated inputs resolve XComArgs + # with ``aresolve``) sees them as well. + mock_supervisor_comms.asend.side_effect = lambda msg, **kwargs: mock_supervisor_comms.send( + msg=msg, **kwargs + ) ti = create_runtime_ti(dag.task_dict[task_id], map_index=map_index) run(ti, ti.get_template_context(), log) @@ -47,3 +54,15 @@ def run(dag: DAG, task_id: str, map_index: int): raise RuntimeError("Unable to find call to TaskState") return run + + +def make_xcom_arg(values: Any) -> XComArg: + op = BaseOperator(task_id="upstream") + xcom_arg = XComArg(op) + xcom_arg.resolve = lambda *a, **kw: values + + async def aresolve(*a, **kw): + return values + + xcom_arg.aresolve = aresolve + return xcom_arg diff --git a/task-sdk/tests/task_sdk/definitions/test_context.py b/task-sdk/tests/task_sdk/definitions/test_context.py index dc25ec378bdee..eb6ae95c4ba3d 100644 --- a/task-sdk/tests/task_sdk/definitions/test_context.py +++ b/task-sdk/tests/task_sdk/definitions/test_context.py @@ -17,9 +17,13 @@ # under the License. from __future__ import annotations +from types import SimpleNamespace + import pytest -from airflow.sdk.definitions.context import get_current_context +from airflow.sdk import Asset +from airflow.sdk.definitions.context import Context, clone_context, get_current_context +from airflow.sdk.execution_time.context import InletEventsAccessors class TestCurrentContext: @@ -35,7 +39,68 @@ def test_get_current_context_with_context(self, monkeypatch): result = get_current_context() assert result == mock_context - def test_get_current_context_without_context(self, monkeypatch): - monkeypatch.setattr("airflow.sdk.definitions._internal.contextmanager._CURRENT_CONTEXT", []) - with pytest.raises(RuntimeError, match="Current context was requested but no context was found!"): - get_current_context() + def test_clone_context_deep_and_shallow_copy_semantics(self): + outlet_events = [Asset(name="dummy")] + inlet_events = InletEventsAccessors(inlets=[]) + + dag_run = SimpleNamespace( + dag_id="dag", + run_id="r1", + logical_date=None, + data_interval_start=None, + data_interval_end=None, + run_after=None, + start_date=None, + end_date=None, + clear_number=None, + run_type=None, + state=None, + conf=None, + triggering_user_name=None, + consumed_asset_events=[], + partition_key=None, + note=None, + ) + + actual = Context() + actual.update( + { + "params": {"p": {"n": 1}}, + "templates_dict": {"tpl": ["a", {"x": 1}]}, + "inlets": [object()], + "outlets": [object()], + "outlet_events": outlet_events, + "inlet_events": inlet_events, + "dag_run": dag_run, + } + ) + cloned = clone_context(actual) + + assert cloned is not actual + + actual["params"]["p"]["n"] = 999 + assert cloned["params"]["p"]["n"] == 1 + + actual["templates_dict"]["tpl"][1]["x"] = 42 + assert cloned["templates_dict"]["tpl"][1]["x"] == 1 + + actual["outlet_events"].append(Asset(name="another")) + assert cloned["outlet_events"] is actual["outlet_events"] + assert len(cloned["outlet_events"]) == 2 + assert Asset(name="another") in cloned["outlet_events"] + assert cloned["inlet_events"] is actual["inlet_events"] + + actual["dag_run"].dag_id = "changed" + assert cloned["dag_run"].dag_id == "changed" + + def test_clone_context_without_inlet_events_or_dag_run(self): + """Context is total=False: a context built outside the runner may lack both keys and still clones.""" + actual = Context() + actual.update({"params": {"p": 1}, "outlet_events": [Asset(name="a")]}) + + cloned = clone_context(actual) + + assert cloned["params"] == {"p": 1} + assert cloned["outlet_events"] is actual["outlet_events"] + assert "inlet_events" not in cloned + assert "dag_run" not in cloned diff --git a/task-sdk/tests/task_sdk/definitions/test_iterableoperator.py b/task-sdk/tests/task_sdk/definitions/test_iterableoperator.py new file mode 100644 index 0000000000000..4ac4e03bc22cf --- /dev/null +++ b/task-sdk/tests/task_sdk/definitions/test_iterableoperator.py @@ -0,0 +1,3933 @@ +# +# 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 + +import asyncio +import threading +import time +from collections.abc import Iterator +from contextlib import contextmanager +from datetime import UTC, datetime, timedelta +from types import SimpleNamespace +from typing import TYPE_CHECKING, Any +from unittest.mock import AsyncMock, MagicMock, Mock, call, create_autospec, patch + +try: + # Python 3.11+ + BaseExceptionGroup +except NameError: + from exceptiongroup import BaseExceptionGroup + +import pytest +from task_sdk.definitions.conftest import make_xcom_arg + +from airflow.sdk import ( + DAG, + Asset, + BaseAsyncOperator, + BaseOperator, + BaseXCom, + TaskInstanceState, + get_current_context, +) +from airflow.sdk.bases.operator import event_loop +from airflow.sdk.definitions._internal.abstractoperator import DEFAULT_RETRIES +from airflow.sdk.definitions._internal.expandinput import ( + DictOfListsExpandInput, + ExpandInput, + ListOfDictsExpandInput, + Resolved, +) +from airflow.sdk.definitions.context import clone_context +from airflow.sdk.definitions.iterableoperator import ( + Checkpoints, + IndexedTaskInstanceNotStarted, + IndexedTaskOutcomes, + IterableOperator, + IterationState, +) +from airflow.sdk.exceptions import ( + AirflowFailException, + AirflowRescheduleException, + AirflowSkipException, + AirflowTaskTerminated, + AirflowTaskTimeout, + DagRunTriggerException, + DownstreamTasksSkipped, + TaskDeferred, + UnmappableXComTypePushed, +) +from airflow.sdk.execution_time.comms import DeadlockImminentError +from airflow.sdk.execution_time.context import InletEventsAccessors, OutletEventAccessors +from airflow.sdk.execution_time.executor import AsyncAwareExecutor +from airflow.sdk.execution_time.task_runner import ( + IndexedTaskInstance, + IndexedTaskRunner, + IndexedTaskState, + RuntimeTaskInstance, +) +from airflow.sdk.execution_time.xcom import XCom + +from tests_common.test_utils.mock_context import mock_context as _mock_context_base + +if TYPE_CHECKING: + from airflow.sdk.definitions._internal.expandinput import ExpandInput + from airflow.sdk.definitions.mappedoperator import MappedOperator + + from tests_common.test_utils.compat import Context + + +class MockTaskStateStoreAccessor: + """Minimal in-memory stand-in for ``TaskStateStoreAccessor``, exposing only the async + ``aget``/``aset`` methods used by ``IterableOperator`` to checkpoint per-index sub-task + progress (see ``IterableOperator._run_task``), plus the sync ``get``/``set``/``delete`` used for + the completion marker (see ``IterableOperator._run_tasks``).""" + + def __init__(self): + self._data: dict[str, Any] = {} + + async def aget(self, key: str, default: Any = None) -> Any: + return self._data.get(key, default) + + async def aset(self, key: str, value: Any, **kwargs) -> None: + self._data[key] = value + + def get(self, key: str, default: Any = None) -> Any: + return self._data.get(key, default) + + def set(self, key: str, value: Any, **kwargs) -> None: + self._data[key] = value + + def delete(self, key: str) -> None: + self._data.pop(key, None) + + def __contains__(self, key: str) -> bool: + return key in self._data + + def __getitem__(self, key: str) -> Any: + return self._data[key] + + +@contextmanager +def mock_context(task, run_id: str | None = None) -> Iterator[Context]: + """Create a mock context for IterableOperator tests. + + The context includes the task state store, asset event accessors, DAG/run + information, and a mocked XCom backend. + """ + task_state_store = MockTaskStateStoreAccessor() + context = _mock_context_base(task=task, run_id=run_id) + context["dag"] = task.dag # type: ignore[typeddict-item] + context["dag_run"] = SimpleNamespace(conf={}) # type: ignore[typeddict-item] + context["task_state_store"] = task_state_store # type: ignore[typeddict-item] + context["outlet_events"] = OutletEventAccessors() + context["inlet_events"] = InletEventsAccessors(inlets=[]) + + def _set( + cls, + key, + value, + *, + dag_id, + task_id, + run_id, + map_index=-1, + **kwargs, + ): + context["ti"].xcom_push(key=key, value=value) + + async def _aset( + cls, + key, + value, + *, + dag_id, + task_id, + run_id, + map_index=-1, + **kwargs, + ): + context["ti"].xcom_push(key=key, value=value) + + def _get_one( + cls, + *, + key, + dag_id, + task_id, + run_id, + map_index=None, + include_prior_dates=False, + **kwargs, + ): + return context["ti"].xcom_pull( + task_ids=task_id, + dag_id=dag_id, + key=key, + ) + + def _get_all(cls, *, key, dag_id, task_id, run_id, include_prior_dates=False, **kwargs): + """One map index in these tests, so the values of all of them are the one value, or none.""" + value = context["ti"].xcom_pull(task_ids=task_id, dag_id=dag_id, key=key) + return None if value is None else [value] + + with ( + patch.object(XCom, "set", classmethod(_set)), + patch.object(XCom, "aset", classmethod(_aset)), + patch.object(XCom, "get_one", classmethod(_get_one)), + patch.object(XCom, "get_all", classmethod(_get_all)), + patch.object(RuntimeTaskInstance, "task_state_store", property(lambda self: task_state_store)), + ): + yield context + + +class MockOperator(BaseOperator): + """Mock operator for testing IterableOperator expansion.""" + + template_fields = ("arg1", "arg2", "arg3") + + def __init__( + self, + arg1=None, + arg2=None, + arg3=None, + fail_on_first_attempt=False, + raise_exception: BaseException | None = None, + **kwargs, + ): + super().__init__(**kwargs) + self.arg1 = arg1 + self.arg2 = arg2 + self.arg3 = arg3 + self.fail_on_first_attempt = fail_on_first_attempt + self.raise_exception = raise_exception + + def execute(self, context): + """Execute the operator and return passed arguments as tuple if do_xcom_push is True.""" + expected = clone_context(context) + + if self.raise_exception is not None: + raise self.raise_exception + if self.fail_on_first_attempt: + self.fail_on_first_attempt = False + raise RuntimeError + if not self.do_xcom_push: + return None + + assert context == expected, "Context was unexpectedly mutated during task execution" + return self.arg1, self.arg2, self.arg3 + + +class MockOperatorWithCustomName(MockOperator): + """MockOperator subclass with a custom display name, mimicking a @task-decorated callable, + used to verify IterableOperator.operator_name forwards the wrapped operator's own + operator_name rather than falling back to this wrapper's task_type.""" + + custom_operator_name = "@mock_task" + + +class MockOutletEventOperator(BaseOperator): + """Operator that records an outlet asset event on execute, used to test that + IterableOperator merges/replays per-sub-task outlet events (see ``_run_task``).""" + + template_fields = () + + def __init__(self, extra_value: str = "v", **kwargs): + super().__init__(**kwargs) + self.extra_value = extra_value + + def execute(self, context): + context["outlet_events"][Asset(name="a", uri="s3://bucket/a")].extra["value"] = self.extra_value + return "done" + + +class MockOnKillOperator(BaseOperator): + """Operator that records whether ``on_kill()`` was called on it, used to test that + IterableOperator.on_kill() propagates to currently in-flight sub-tasks (see ``on_kill``).""" + + template_fields = () + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.killed = False + + def execute(self, context): + return "done" + + def on_kill(self): + self.killed = True + + +class MockDeferredOperator(BaseOperator): + """Operator that immediately defers on execute, simulating a deferrable operator.""" + + template_fields = () + + def execute(self, context): + raise TaskDeferred(trigger=None, method_name="execute_complete") # type: ignore[arg-type] + + +class MockRescheduleSensor(BaseOperator): + """Operator that raises AirflowRescheduleException on execute, simulating a reschedule-mode sensor.""" + + template_fields = () + + def execute(self, context): + from datetime import timedelta + + from airflow.sdk import timezone + + raise AirflowRescheduleException(timezone.utcnow() + timedelta(seconds=60)) + + +class MockStateStoreOperator(BaseOperator): + """Sync operator that keeps its own state in the task state store, as a paginated fetch would.""" + + template_fields = ("offset",) + + def __init__(self, offset=None, **kwargs): + super().__init__(**kwargs) + self.offset = offset + + def execute(self, context): + store = context["task_state_store"] + assert store is context["ti"].task_state_store + store.set("last_offset", self.offset) + return store.get("last_offset") + + +class MockAsyncStateStoreOperator(BaseAsyncOperator): + """Async twin of MockStateStoreOperator.""" + + template_fields = ("offset",) + + def __init__(self, offset=None, **kwargs): + super().__init__(**kwargs) + self.offset = offset + + async def aexecute(self, context): + store = context["task_state_store"] + await store.aset("last_offset", self.offset) + return await store.aget("last_offset") + + +class MockAttemptOperator(BaseAsyncOperator): + """ + Async operator whose result tells which attempt produced it. + + For the values in ``times_out_on`` the parent's execution timeout strikes instead, which ends + the whole iteration there and leaves the items after it unreached, as a crash would. Async, + because that timeout lands on the main thread, where async operators run. + """ + + template_fields = ("arg1",) + times_out_on: set = set() + + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + async def aexecute(self, context): + if self.arg1 in self.times_out_on: + raise AirflowTaskTimeout("the iteration ran out of time") + return f"{self.arg1}@attempt{context['ti'].try_number}" + + +FIRED_CALLBACKS: list = [] + + +class MockCallbackTimeoutOperator(BaseAsyncOperator): + """ + Async operator in which the parent's execution timeout strikes; records which callback fired. + + Async, because the parent's limit lands on the main thread, where async operators run; a sync + operator raising ``AirflowTaskTimeout`` in its worker thread raised its own. + """ + + template_fields = () + + def __init__(self, **kwargs): + kwargs["on_failure_callback"] = lambda context: FIRED_CALLBACKS.append("failure") + kwargs["on_retry_callback"] = lambda context: FIRED_CALLBACKS.append("retry") + super().__init__(**kwargs) + + async def aexecute(self, context): + raise AirflowTaskTimeout("the task ran out of time") + + +class MockCallbackAsyncOperator(BaseAsyncOperator): + """Async operator that defers for ``arg1="defer"`` and sleeps otherwise; records its callbacks.""" + + template_fields = ("arg1",) + + def __init__(self, arg1=None, **kwargs): + kwargs["on_failure_callback"] = lambda context: FIRED_CALLBACKS.append(("failure", self.arg1)) + kwargs["on_retry_callback"] = lambda context: FIRED_CALLBACKS.append(("retry", self.arg1)) + super().__init__(**kwargs) + self.arg1 = arg1 + + async def aexecute(self, context): + if self.arg1 == "defer": + raise TaskDeferred(trigger=None, method_name="execute_complete") # type: ignore[arg-type] + await asyncio.sleep(60) + + +class MockPushingOperator(BaseOperator): + """Operator that pushes an extra XCom next to its return value, and fails for ``fail=True``.""" + + template_fields = ("arg1",) + executed: list = [] + + def __init__(self, arg1=None, fail=False, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + self.fail = fail + + def execute(self, context): + type(self).executed.append(self.arg1) + context["ti"].xcom_push(key="foo", value=f"foo-of-{self.arg1}") + if self.fail: + raise RuntimeError("sibling failed") + return self.arg1 + + +class MockClearingStateStoreOperator(BaseOperator): + template_fields = () + + def execute(self, context): + context["task_state_store"].clear() + + +def wait_until(condition, timeout: float = 5.0) -> bool: + """Poll ``condition`` until it holds; on_kill() kills in a thread of its own.""" + deadline = time.monotonic() + timeout + while not condition(): + if time.monotonic() > deadline: + return False + time.sleep(0.01) + return True + + +def create_mapped_operator( + dag: DAG, + expand_input: ExpandInput, + task_id: str = "my_task", + retries: int = DEFAULT_RETRIES, + do_xcom_push: bool = True, + task_concurrency: int | None = None, + execution_timeout: timedelta | None = None, + operator_class: type[BaseOperator] = MockOperator, +) -> MappedOperator: + """ + Create a MappedOperator and assign it to a DAG. + + :param expand_input: The input to expand + :param dag: The DAG to assign the operator to + :param task_id: Task ID for the operator + :param do_xcom_push: Whether to push XCom (default True) + :param operator_class: Operator class to wrap (default MockOperator) + """ + return operator_class.partial( + task_id=task_id, + dag=dag, + retries=retries, + task_concurrency=task_concurrency, + do_xcom_push=do_xcom_push, + execution_timeout=execution_timeout, + )._expand( + expand_input, + strict=True, + register_with_dag=False, + ) + + +def create_iterable_operator( + dag: DAG, + expand_input: ExpandInput, + task_id: str = "my_task", + task_concurrency: int | None = None, + retries: int = DEFAULT_RETRIES, + do_xcom_push: bool = True, + operator_class: type[BaseOperator] = MockOperator, +) -> IterableOperator: + """Create an IterableOperator with a MappedOperator and ExpandInput.""" + mapped_op = create_mapped_operator( + dag=dag, + expand_input=expand_input, + task_id=task_id, + retries=retries, + do_xcom_push=do_xcom_push, + task_concurrency=task_concurrency, + operator_class=operator_class, + ) + return IterableOperator( + operator=mapped_op, + expand_input=expand_input, + dag=dag, + ) + + +def _items(expand_input: ExpandInput) -> list: + """Every item of the input in index order, the way IterableOperator.execute reads it.""" + + async def read(): + length, aget = await expand_input.aresolve({}) + return [await aget(index) for index in range(length)] + + return asyncio.run(read()) + + +class TestIterableOperator: + @pytest.mark.parametrize( + ("actual", "expected"), + [ + ([{"a": 1}, {"a": 2}], [{"a": 1}, {"a": 2}]), + ([{"a": 1, "b": 2}], [{"a": 1, "b": 2}]), + ([], []), + ], + ) + def test_list_of_dicts_expand_input_aresolve(self, actual, expected): + """Test IterableOperator with ListOfDictsExpandInput expand_input.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput(actual) + iterable_op = create_iterable_operator(dag, expand_input) + + assert _items(iterable_op.expand_input) == expected + + @pytest.mark.parametrize( + ("actual", "expected"), + [ + ({"a": [1, 2, 3]}, [{"a": 1}, {"a": 2}, {"a": 3}]), + ( + {"a": [1, 2], "b": [10, 20]}, + [{"a": 1, "b": 10}, {"a": 1, "b": 20}, {"a": 2, "b": 10}, {"a": 2, "b": 20}], + ), + ({"a": [1, 2]}, [{"a": 1}, {"a": 2}]), + ( + {"a": {"x": 1, "y": 2}}, + [{"a": ("x", 1)}, {"a": ("y", 2)}], + ), + ], + ) + def test_dict_of_lists_expand_input_aresolve(self, actual, expected): + """Test IterableOperator with DictOfListsExpandInput expand_input. + + A dict value expands to its (key, value) pairs (not just its keys), matching + the classic .expand() resolve() path's handling of dict values. + """ + with DAG("test_dag") as dag: + expand_input = DictOfListsExpandInput(actual) + iterable_op = create_iterable_operator(dag, expand_input) + + assert _items(iterable_op.expand_input) == expected + + def test_task_type(self): + """ + IterableOperator reports its own class as its type, so whatever resolves a class by it + (OpenLineage's extractors, the task's class reference) finds the class that runs, while + the name it is shown under is the wrapped operator's. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + iterable_op = create_iterable_operator(dag, expand_input) + + assert isinstance(iterable_op, IterableOperator) + assert iterable_op.task_type == "IterableOperator" + assert iterable_op.operator_name == "MockOperator" + + def test_operator_name(self): + """Test that IterableOperator forwards the wrapped operator's operator_name (e.g. a + @task-decorated callable's custom_operator_name), not just its own task_type.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + iterable_op = create_iterable_operator( + dag, expand_input, operator_class=MockOperatorWithCustomName + ) + + assert isinstance(iterable_op, IterableOperator) + assert iterable_op.task_type == "IterableOperator" + assert iterable_op.operator_name == "@mock_task" + + def test_forwards_params_weight_rule_and_retry_policy(self): + """Test that IterableOperator forwards params, weight_rule, and retry_policy from the + wrapped operator onto its own DAG node, not just onto the generated sub-tasks.""" + from airflow.sdk import WeightRule + from airflow.sdk.definitions.retry_policy import ExceptionRetryPolicy + + retry_policy = ExceptionRetryPolicy(rules=[]) + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + mapped_op = MockOperator.partial( + task_id="my_task", + dag=dag, + params={"p": 1}, + weight_rule=WeightRule.UPSTREAM, + retry_policy=retry_policy, + )._expand(expand_input, strict=True, register_with_dag=False) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + assert iterable_op.params["p"] == 1 + assert iterable_op.weight_rule == WeightRule.UPSTREAM + assert iterable_op.retry_policy is retry_policy + + def test_default_args_callbacks_and_hooks_stay_with_the_items(self): + """ + Callbacks and execute hooks in ``default_args`` reach the items, not the iterated task. + + ``_apply_defaults`` fills every ``BaseOperator.__init__`` parameter the call leaves out + from ``default_args``; the IterableOperator passes them as ``None``, so a DAG-level + ``on_failure_callback`` runs once per failed item, as through ``.partial()``, and not once + more for the task with the parent's context. + """ + callback = Mock() + with DAG( + "test_dag", + default_args={ + "on_failure_callback": callback, + "on_success_callback": callback, + "on_execute_callback": callback, + "pre_execute": callback, + "post_execute": callback, + }, + ): + iterated = MockOperator.partial(task_id="op").iterate(arg1=["a", "b"]) + + assert iterated.on_failure_callback == [] + assert iterated.on_success_callback == [] + assert iterated.on_execute_callback == [] + assert iterated._pre_execute_hook is None + assert iterated._post_execute_hook is None + item = iterated._operator.unmap({"arg1": "a"}) + assert item.on_failure_callback == [callback] + assert item.on_success_callback == [callback] + assert item._pre_execute_hook is callback + + def test_forwards_do_xcom_push(self): + """Test that IterableOperator forwards do_xcom_push from the wrapped operator's + partial_kwargs onto its own DAG node.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + iterable_op = create_iterable_operator(dag, expand_input, do_xcom_push=False) + + assert iterable_op.do_xcom_push is False + + def test_forwards_is_setup_is_teardown_and_on_failure_fail_dagrun(self): + """Test that IterableOperator forwards is_setup/is_teardown/on_failure_fail_dagrun from + the wrapped operator's partial_kwargs, mirroring what unmap() reads (via + ``_get_unmap_kwargs``) to apply the same flags to each generated sub-task.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + mapped_op = create_mapped_operator(dag, expand_input) + # unmap() only ever applies these three flags to sub-tasks by reading them off + # partial_kwargs (see MappedOperator._get_unmap_kwargs), so that's the only place + # IterableOperator needs to source them from too. + mapped_op.partial_kwargs["is_teardown"] = True + mapped_op.partial_kwargs["on_failure_fail_dagrun"] = True + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + assert iterable_op.is_setup is False + assert iterable_op.is_teardown is True + assert iterable_op.on_failure_fail_dagrun is True + + def test_applies_upstream_relationship_for_partial_kwargs_template_fields(self): + """Test that an XComArg passed via a partial kwarg matching the wrapped operator's own + template_fields is wired as an upstream dependency of the IterableOperator, mirroring what + MappedOperator.__attrs_post_init__ does for a normal mapped task.""" + with DAG("test_dag") as dag: + upstream = MockOperator(task_id="upstream", dag=dag) + expand_input = ListOfDictsExpandInput([{"arg1": 1}]) + mapped_op = MockOperator.partial( + task_id="my_task", + dag=dag, + arg1=upstream.output, + )._expand(expand_input, strict=True, register_with_dag=False) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + assert upstream.task_id in iterable_op.upstream_task_ids + + def test_iterate_inside_mapped_task_group_raises_not_implemented_error(self): + """Test that wrapping an operator expanded inside a mapped task group with IterableOperator + raises NotImplementedError, since operator expansion in an expanded task group is not + supported (mirrors MappedOperator's own guard for the analogous .expand() case).""" + from airflow.decorators import task_group + + with DAG("test_dag") as dag: + + @task_group + def tg(va): + expand_input = ListOfDictsExpandInput([{"arg1": 1}]) + mapped_op = MockOperator.partial( + task_id="my_task", + dag=dag, + )._expand(expand_input, strict=True, register_with_dag=False) + + with pytest.raises(NotImplementedError, match="operator expansion in an expanded task group"): + IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + tg.expand(va=[["a", "b"], [4]]) + + def test_task_retries(self): + """Test that IterableOperator inherits retries from the wrapped operator, since + the whole IterableOperator is now retried via Airflow's standard retry mechanism.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + iterable_op = create_iterable_operator(dag, expand_input, retries=3) + + assert isinstance(iterable_op, IterableOperator) + assert iterable_op.retries == 3 + assert iterable_op.task_retries == 3 + + def test_task_id(self): + """Test that IterableOperator inherits task_id from operator.""" + with DAG("test_dag") as dag: + task_id = "my_task" + expand_input = ListOfDictsExpandInput([{"a": 1}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id=task_id) + + assert iterable_op.task_id == task_id + + def test_with_task_concurrency(self): + """Test that IterableOperator respects task_concurrency parameter.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + iterable_op = create_iterable_operator(dag, expand_input, task_concurrency=4) + + assert iterable_op.max_workers == 4 + + def test_direct_instantiation_rejects_task_concurrency(self): + """A directly instantiated operator can never reach IterableOperator, + so task_concurrency must be rejected instead of silently accepted as a dead value.""" + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + MockOperator(task_id="my_task", task_concurrency=4) + + def test_expand_rejects_task_concurrency(self): + """.expand() produces a plain MappedOperator, never an IterableOperator, so + task_concurrency (only meaningful for .iterate()/.iterate_kwargs()) must be rejected.""" + with DAG("test_dag"): + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + MockOperator.partial(task_id="my_task", task_concurrency=4).expand(arg1=[1, 2, 3]) + + def test_expand_kwargs_rejects_task_concurrency(self): + """.expand_kwargs() produces a plain MappedOperator, never an IterableOperator, so + task_concurrency (only meaningful for .iterate()/.iterate_kwargs()) must be rejected.""" + with DAG("test_dag"): + with pytest.raises(TypeError, match="which is now max_active_tis_per_dag"): + MockOperator.partial(task_id="my_task", task_concurrency=4).expand_kwargs( + [{"arg1": 1}, {"arg1": 2}] + ) + + def test_iterate_accepts_task_concurrency(self): + """.iterate() is the one public entry point where task_concurrency is meaningful: it + produces an IterableOperator, which reads task_concurrency out of partial_kwargs as + max_workers rather than forwarding it to BaseOperator.__init__.""" + with DAG("test_dag"): + iterable_op = MockOperator.partial(task_id="my_task", task_concurrency=4).iterate(arg1=[1, 2, 3]) + + assert isinstance(iterable_op, IterableOperator) + assert iterable_op.max_workers == 4 + + def test_iterate_kwargs_accepts_task_concurrency(self): + """.iterate_kwargs() is the list-of-dicts counterpart to .iterate() and must accept + task_concurrency the same way.""" + with DAG("test_dag"): + iterable_op = MockOperator.partial(task_id="my_task", task_concurrency=4).iterate_kwargs( + [{"arg1": 1}, {"arg1": 2}] + ) + + assert isinstance(iterable_op, IterableOperator) + assert iterable_op.max_workers == 4 + + @pytest.mark.parametrize("invalid_value", [0, -1, -10]) + def test_task_concurrency_validation_rejects_non_positive_values(self, invalid_value): + """Test that IterableOperator raises ValueError for task_concurrency < 1.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + with pytest.raises(ValueError, match=f"task_concurrency must be at least 1, got {invalid_value}"): + create_iterable_operator(dag, expand_input, task_concurrency=invalid_value) + + def test_partial_kwargs_not_mutated(self): + """Test that creating IterableOperator does not mutate the original MappedOperator's partial_kwargs.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + mapped_op = create_mapped_operator(dag, expand_input, task_concurrency=4) + original_partial_kwargs = mapped_op.partial_kwargs.copy() + + IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + # Verify that mapped_op.partial_kwargs was not mutated + assert mapped_op.partial_kwargs == original_partial_kwargs + assert "task_concurrency" in mapped_op.partial_kwargs + + def test_expand_input_stored(self): + """Test that IterableOperator stores expand_input correctly.""" + with DAG("test_dag") as dag: + expand_input_data = ListOfDictsExpandInput([{"a": 1}, {"a": 2}]) + iterable_op = create_iterable_operator(dag, expand_input_data) + + assert iterable_op.expand_input is expand_input_data + assert isinstance(iterable_op.expand_input, (ListOfDictsExpandInput, DictOfListsExpandInput)) + + def test_partial_kwargs_stored(self): + """Test that IterableOperator stores partial_kwargs from operator.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"a": 1}]) + iterable_op = create_iterable_operator(dag, expand_input) + + assert hasattr(iterable_op, "partial_kwargs") + assert isinstance(iterable_op.partial_kwargs, dict) + + def test_xcom_push_delegates_to_task(self): + """_xcom_push awaits task.axcom_push with the default XCom return key.""" + from unittest import mock + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator(dag, expand_input) + + task = mock.MagicMock() + task.axcom_push = mock.AsyncMock() + task.task_id = "my_task" + task.index = 0 + + asyncio.run(iterable_op.axcom_push(task=task, value="result_value")) + + task.axcom_push.assert_awaited_once_with(key=BaseXCom.XCOM_RETURN_KEY, value="result_value") + + def test_execute_list_of_dicts(self): + """Test executing IterableOperator with ListOfDictsExpandInput.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_list_of_dicts") + + with mock_context(task=iterable_op) as context: + result = iterable_op.execute(context=context) + materialized = list(result) + assert materialized == [(1, None, None), (2, None, None)] + + @pytest.mark.parametrize( + "operator_class", + [MockStateStoreOperator, MockAsyncStateStoreOperator], + ids=["sync", "async"], + ) + def test_execute_gives_each_iteration_its_own_task_state_store_keys(self, operator_class): + """Iterations share one task instance, so a key an iteration stores carries its index.""" + with DAG("test_dag") as dag: + expand_input = DictOfListsExpandInput({"offset": [10, 20, 30]}) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="exec_state_store", operator_class=operator_class + ) + + with mock_context(task=iterable_op) as context: + results = list(iterable_op.execute(context=context)) + store = context["task_state_store"] + + assert results == [10, 20, 30] + assert {key: store[key] for key in ("last_offset_0", "last_offset_1", "last_offset_2")} == { + "last_offset_0": 10, + "last_offset_1": 20, + "last_offset_2": 30, + } + assert "last_offset" not in store + # the operator's own checkpoints keep their unsuffixed keys in the same store + assert all(f"_iterable_{index}" in store for index in range(3)) + + def test_task_state_store_clear_is_refused_inside_an_iteration(self): + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="exec_clear", operator_class=MockClearingStateStoreOperator + ) + + with mock_context(task=iterable_op) as context: + with pytest.raises(RuntimeError, match="not available inside an iterated task"): + iterable_op.execute(context=context) + + def test_execute_marks_iteration_completed_once_every_index_succeeds(self): + """ + Once every sub-task index has succeeded, ``_run_tasks`` writes a single completion marker + instead of deleting the per-index checkpoints one by one, and leaves any state written by + user code inside the iterated task untouched. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_completed") + + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + store._data["last_offset"] = 42 + + list(iterable_op.execute(context=context)) + + assert store["last_offset"] == 42 + assert store["_iterable_completed"]["completed"] is True + assert store["_iterable_0"]["status"] == "success" + assert store["_iterable_1"]["status"] == "success" + + def test_execute_reruns_every_index_after_a_clear_that_follows_success(self): + """ + A manual clear does not reset the parent TI's ``try_number`` (see ``clear_task_instances``), + so the next attempt looks like a retry. The completion marker left by the previous fully + successful run tells ``_run_tasks`` to ignore the stale ``SUCCESS`` checkpoints and run every + index again, recording first where the rerun starts so a crash mid-rerun resumes from the + new checkpoints only. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_rerun") + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 2 + store = context["task_state_store"] + store._data["_iterable_completed"] = {"completed": True, "try_number": 1} + for index in (0, 1): + store._data[f"_iterable_{index}"] = IndexedTaskState( + status=TaskInstanceState.SUCCESS, result="stale" + ).serialize() + + materialized = list(iterable_op.execute(context=context)) + + assert materialized == [(1, None, None), (2, None, None)] + assert IndexedTaskState.deserialize(store["_iterable_0"]).result == (1, None, None) + assert store["_iterable_completed"] == {"completed": True, "try_number": 2} + + def test_execute_after_a_crashed_rerun_does_not_replay_results_from_before_the_clear(self, monkeypatch): + """ + A task that succeeded is cleared and its rerun stops part-way. The next attempt resumes + what the rerun finished and runs the rest again: an index the rerun did not get to still + holds its checkpoint from before the clear, and replaying that one is not what clearing + the task asked for. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}, {"arg1": 3}]) + iterable_op = create_iterable_operator( + dag, + expand_input, + task_id="exec_crashed_rerun", + operator_class=MockAttemptOperator, + task_concurrency=1, + ) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + first_run = list(iterable_op.execute(context=context)) + assert first_run == ["1@attempt1", "2@attempt1", "3@attempt1"] + + # Cleared; the rerun finishes the first item and stops at the second. + context["ti"].try_number = 2 + monkeypatch.setattr(MockAttemptOperator, "times_out_on", {2}) + with pytest.raises(AirflowTaskTimeout): + iterable_op.execute(context=context) + + context["ti"].try_number = 3 + monkeypatch.setattr(MockAttemptOperator, "times_out_on", set()) + materialized = list(iterable_op.execute(context=context)) + + assert materialized == ["1@attempt2", "2@attempt3", "3@attempt3"] + assert context["task_state_store"]["_iterable_completed"] == { + "completed": True, + "try_number": 3, + } + + def test_execute_resumes_from_checkpoints_on_retry_after_failure(self): + """ + On a genuine retry (no completion marker) an index checkpointed as ``SUCCESS`` is skipped and + its checkpointed result is pushed again, while an index left ``UP_FOR_RETRY`` runs, and the + returned iterable yields both. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_resume") + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 2 + store = context["task_state_store"] + store._data["_iterable_0"] = IndexedTaskState( + status=TaskInstanceState.SUCCESS, + result="from_checkpoint", + fingerprint=IterableOperator._fingerprint({"arg1": 1}), + ).serialize() + store._data["_iterable_1"] = IndexedTaskState( + status=TaskInstanceState.UP_FOR_RETRY + ).serialize() + + materialized = list(iterable_op.execute(context=context)) + + assert materialized == ["from_checkpoint", (2, None, None)] + assert store["_iterable_1"]["status"] == "success" + assert store["_iterable_completed"]["completed"] is True + + def test_execute_runs_an_index_again_when_its_input_changed(self): + """ + A retry can run on another input than the attempt that wrote the checkpoints: the upstream + was cleared together with this task and produced other items. The index whose item changed + runs again instead of replaying a result computed from the old item; the other is resumed. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": "new"}, {"arg1": 2}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_input_changed") + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 2 + store = context["task_state_store"] + store._data["_iterable_0"] = IndexedTaskState( + status=TaskInstanceState.SUCCESS, + result="from_the_old_item", + fingerprint=IterableOperator._fingerprint({"arg1": "old"}), + ).serialize() + store._data["_iterable_1"] = IndexedTaskState( + status=TaskInstanceState.SUCCESS, + result="from_checkpoint", + fingerprint=IterableOperator._fingerprint({"arg1": 2}), + ).serialize() + + materialized = list(iterable_op.execute(context=context)) + + assert materialized == [("new", None, None), "from_checkpoint"] + assert store["_iterable_0"]["fingerprint"] == IterableOperator._fingerprint({"arg1": "new"}) + + def test_fingerprint(self): + """The digest follows the content, not the key order, and gives up on what serde cannot serialize.""" + assert IterableOperator._fingerprint({"a": 1, "b": [1, 2]}) == IterableOperator._fingerprint( + {"b": [1, 2], "a": 1} + ) + assert IterableOperator._fingerprint({"a": 1}) != IterableOperator._fingerprint({"a": 2}) + assert IterableOperator._fingerprint({"a": object()}) is None + + def test_execute_does_not_leak_unmapped_operator_into_parent_context(self): + """ + Regression test: unmapping a sub-task must not mutate the parent context's own `ti`. + + ``context_update_for_unmapped`` sets ``context["ti"].task = task`` in place. Since + ``context.copy()`` is only a shallow copy, ``context["ti"]`` in the copy is the *same* + object as the parent's. If ``_create_task`` rendered the unmapped sub-operator against a + context still carrying the parent's `ti`, the parent's `ti.task` would end up pointing at + whichever sub-task was unmapped last, corrupting anything the runner reads off `ti.task` + after `execute()` returns (e.g. `do_xcom_push`/`multiple_outputs` in `_push_xcom_if_needed`). + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}, {"arg1": 3}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_no_leak") + + with mock_context(task=iterable_op) as context: + parent_ti_task = context["ti"].task + + list(iterable_op.execute(context=context)) + + assert context["ti"].task is parent_ti_task + assert context["task"] is iterable_op + + def test_an_item_reads_back_the_xcom_it_pushed(self): + """ + ``ti.xcom_push("progress", v)`` in an item lands under ``progress_``; the item's own + ``ti.xcom_pull(key="progress")`` must read that back, not the parent's unsuffixed key. + """ + seen: list = [] + + class PushThenPull(BaseOperator): + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + def execute(self, context): + ti = context["ti"] + ti.xcom_push(key="progress", value=self.arg1 * 10) + seen.append( + ( + self.arg1, + ti.xcom_pull(key="progress"), + ti.xcom_pull(task_ids=ti.task_id, key="progress"), + ) + ) + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + mapped_op = PushThenPull.partial(task_id="push_pull", dag=dag, task_concurrency=1)._expand( + expand_input, strict=True, register_with_dag=False + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + with mock_context(task=iterable_op) as context: + iterable_op.execute(context=context) + + assert sorted(seen) == [(1, 10, 10), (2, 20, 20)] + + def test_an_item_sees_its_own_operator_as_the_contexts_task(self): + """ + ``context["task"]`` inside an item, and in its callbacks, is the item's unmapped operator. + + ``.expand()`` gives each mapped task instance its own operator under ``context["task"]`` + (``context_update_for_unmapped`` sets it next to ``ti.task``); the item's view of the context + must swap that key as well, not only ``ti``, or user code reads the IterableOperator there. + """ + seen: list[tuple[str, object, object]] = [] + + class MockContextTaskOperator(BaseOperator): + def execute(self, context): + seen.append(("execute", context["task"], context["ti"].task)) + + def on_success(context): + seen.append(("on_success_callback", context["task"], context["ti"].task)) + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}, {}]) + mapped_op = MockContextTaskOperator.partial( + task_id="context_task", dag=dag, on_success_callback=on_success + )._expand(expand_input, strict=True, register_with_dag=False) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + iterable_op.execute(context=context) + + assert len(seen) == 4 + for where, task, ti_task in seen: + assert task is ti_task, where + assert isinstance(task, MockContextTaskOperator), where + assert task is not iterable_op, where + assert context["task"] is iterable_op + + def test_execute_resolves_and_reads_the_expand_input_on_the_running_loop(self): + """ + Regression test for the frozen IterableOperator: sub-task inputs were pulled from the main + thread between two runs of the event loop, where a blocking supervisor call deadlocked with + the ``asend`` of a sub-task parked mid-call. The input must be resolved with ``aresolve`` + and every index read while the loop runs; the synchronous ``resolve`` must stay untouched. + """ + loop_running: list[bool] = [] + original_aresolve = ListOfDictsExpandInput.aresolve + + async def aresolve(self, context): + loop_running.append(asyncio.get_running_loop().is_running()) + length, aget = await original_aresolve(self, context) + + async def aget_on_loop(index): + loop_running.append(asyncio.get_running_loop().is_running()) + return await aget(index) + + return Resolved(length, aget_on_loop) + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}, {"arg1": 3}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_async_input") + + with ( + mock_context(task=iterable_op) as context, + patch.object(ListOfDictsExpandInput, "aresolve", aresolve), + patch.object( + ListOfDictsExpandInput, "resolve", side_effect=AssertionError("sync resolve used") + ), + ): + materialized = sorted(iterable_op.execute(context=context)) + + assert materialized == [(1, None, None), (2, None, None), (3, None, None)] + assert loop_running == [True, True, True, True] + + def test_execute_pulls_xcom_arg_inputs_through_aresolve(self): + """An XComArg input is resolved with ``aresolve`` (``ti.axcom_pull``), never with blocking ``resolve``.""" + with DAG("test_dag") as dag: + xcom_arg = make_xcom_arg([{"arg1": 1}, {"arg1": 2}]) + xcom_arg.resolve = lambda *a, **kw: pytest.fail("synchronous resolve() used on the loop") + expand_input = ListOfDictsExpandInput(xcom_arg) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_xcom_arg_input") + + with mock_context(task=iterable_op) as context: + materialized = sorted(iterable_op.execute(context=context)) + + assert materialized == [(1, None, None), (2, None, None)] + + def test_execute_refuses_an_unmappable_upstream_value(self): + """ + An upstream returning a JSON string fails the task instead of being iterated as one item. + + ``.expand()`` gets this from the upstream's push (``UnmappableXComTypePushed`` in + ``_push_xcom_if_needed``), which only fires for a ``MappedOperator`` dependant, so the + iterated path checks the resolved value itself. + """ + executed: list = [] + + class RecordingOperator(MockOperator): + def execute(self, context): + executed.append(self.arg1) + return super().execute(context) + + with DAG("test_dag") as dag: + expand_input = DictOfListsExpandInput({"arg1": make_xcom_arg('{"a": 1}')}) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="exec_unmappable", operator_class=RecordingOperator + ) + with mock_context(task=iterable_op) as context: + with pytest.raises(UnmappableXComTypePushed, match="str"): + iterable_op.execute(context=context) + + assert executed == [] + + def test_a_template_reads_the_indexed_tasks_own_state_store(self): + """A template field reads the same suffixed store as ``execute`` does, not the parent's key.""" + template = "{{ task_state_store.get('offset') }}" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"offset": template}, {"offset": template}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="template_store", operator_class=MockStateStoreOperator + ) + + with mock_context(task=iterable_op) as context: + context["task_state_store"].set("offset", "parents") + context["task_state_store"].set("offset_1", 7) + task = iterable_op._create_task( + context=context, + index=1, + mapped_kwargs={"offset": template}, + jinja_env=iterable_op.get_template_env(dag=dag), + ) + + assert task.task.offset == "7" + + def test_a_template_reads_the_items_own_index(self): + """ + ``{{ ti.index }}`` in a partial kwarg renders per item, so an operator that finds its remote + work by labels (a ``KubernetesPodOperator`` that reattaches) can be told the item apart. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": "a"}, {"arg1": "b"}]) + mapped_op = MockOperator.partial(task_id="indexed_label", dag=dag, arg2="{{ ti.index }}")._expand( + expand_input, strict=True, register_with_dag=False + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + results = sorted(iterable_op.execute(context=context)) + + assert results == [("a", "0", None), ("b", "1", None)] + + def test_execute_renders_template_fields_off_the_loop_thread(self): + """ + Rendering a sub-task's template fields may call the supervisor synchronously: an XComArg in + a partial kwarg resolves with ``resolve``, ``{{ var.value.x }}`` with ``Variable.get``. On + the loop thread that call raises ``DeadlockImminentError`` whenever a sibling's ``asend`` is + in flight, so the rendering has to happen in a worker thread. + """ + on_running_loop: list[bool] = [] + + def resolve(context): + try: + asyncio.get_running_loop() + except RuntimeError: + on_running_loop.append(False) + else: + on_running_loop.append(True) + return "pulled" + + with DAG("test_dag") as dag: + xcom_arg = make_xcom_arg(None) + xcom_arg.resolve = resolve + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + mapped_op = MockOperator.partial(task_id="render_off_loop", dag=dag, arg2=xcom_arg)._expand( + expand_input, strict=True, register_with_dag=False + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + materialized = sorted(iterable_op.execute(context=context)) + + assert materialized == [(1, "pulled", None), (2, "pulled", None)] + assert on_running_loop == [False, False] + + def test_execute_dict_of_lists(self): + """Test executing IterableOperator with DictOfListsExpandInput.""" + with DAG("test_dag") as dag: + expand_input = DictOfListsExpandInput({"arg1": [1, 2, 3]}) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_dict_of_lists") + + with mock_context(task=iterable_op) as context: + result = iterable_op.execute(context=context) + materialized = list(result) + assert materialized == [(1, None, None), (2, None, None), (3, None, None)] + + def test_execute_multiple_key_dict_of_lists(self): + """Test executing IterableOperator with multiple keys in DictOfListsExpandInput.""" + with DAG("test_dag") as dag: + expand_input = DictOfListsExpandInput({"arg1": [1, 2], "arg2": [10, 20], "arg3": ["x", "y"]}) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_multi_key") + + with mock_context(task=iterable_op) as context: + result = iterable_op.execute(context=context) + materialized = list(result) + # Cartesian product expected order: + # (1,10,'x'), (1,10,'y'), (1,20,'x'), (1,20,'y'), + # (2,10,'x'), (2,10,'y'), (2,20,'x'), (2,20,'y') + assert materialized == [ + (1, 10, "x"), + (1, 10, "y"), + (1, 20, "x"), + (1, 20, "y"), + (2, 10, "x"), + (2, 10, "y"), + (2, 20, "x"), + (2, 20, "y"), + ] + + def test_execute_with_task_concurrency_setting(self): + """Test executing IterableOperator with task_concurrency parameter.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}, {"arg1": 3}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="exec_concurrency", task_concurrency=2 + ) + + with mock_context(task=iterable_op) as context: + result = iterable_op.execute(context=context) + materialized = list(result) + assert materialized == [(1, None, None), (2, None, None), (3, None, None)] + assert iterable_op.max_workers == 2 + + def test_execute_all_parameters(self): + """Test executing IterableOperator with all arg1, arg2, arg3 parameters.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [ + {"arg1": 1, "arg2": 10, "arg3": 100}, + {"arg1": 2, "arg2": 20, "arg3": 200}, + ] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_all_args") + + with mock_context(task=iterable_op) as context: + result = iterable_op.execute(context=context) + materialized = list(result) + assert materialized == [(1, 10, 100), (2, 20, 200)] + + def test_execute_with_do_xcom_push_false(self): + """With do_xcom_push=False no return_value_ XCom is pushed for any sub-task.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="no_xcom_push", do_xcom_push=False + ) + + with ( + mock_context(task=iterable_op) as context, + patch.object( + IterableOperator, "axcom_push", new=AsyncMock(spec=IterableOperator.axcom_push) + ) as axcom_push, + ): + result = iterable_op.execute(context=context) + + assert result is None + axcom_push.assert_not_awaited() + + def test_execute_does_not_push_xcom_for_none_results(self): + """A sub-task returning None has nothing to push, even when do_xcom_push is True.""" + + class NoneOperator(BaseOperator): + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + def execute(self, context): + return None + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="none_results", operator_class=NoneOperator + ) + + with ( + mock_context(task=iterable_op) as context, + patch.object( + IterableOperator, "axcom_push", new=AsyncMock(spec=IterableOperator.axcom_push) + ) as axcom_push, + ): + iterable_op.execute(context=context) + + axcom_push.assert_not_awaited() + + def test_execute_with_failed_tasks_raises_regardless_of_retries(self): + """ + Test executing IterableOperator where a sub-task fails. + + This test verifies that: + 1. Tasks with fail_on_first_attempt=True raise an exception on first attempt + 2. IterableOperator no longer retries failed sub-tasks in-process — retries (if any) are + handled by Airflow retrying the whole IterableOperator task instance + 3. The failing sub-task's own exception is raised, regardless of whether the wrapped operator + has retries configured (one failure is not wrapped in a group, see _failure_for_the_runner) + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [ + {"arg1": 1, "arg2": 10}, + {"arg1": 2, "arg2": 20, "fail_on_first_attempt": True}, + {"arg1": 3, "arg2": 30}, + ] + ) + iterable_op = create_iterable_operator( + dag, + expand_input, + task_id="exec_with_failures", + retries=1, + ) + + with mock_context(task=iterable_op) as context: + with pytest.raises(RuntimeError): + iterable_op.execute(context=context) + + def test_execute_all_sub_tasks_skipped_raises_single_skip_exception(self): + """When every sub-task raises AirflowSkipException, IterableOperator must re-raise a single + AirflowSkipException so the runner marks it SKIPPED instead of a retryable failure.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [ + {"raise_exception": AirflowSkipException("skip 1")}, + {"raise_exception": AirflowSkipException("skip 2")}, + ] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="all_skipped") + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowSkipException): + iterable_op.execute(context=context) + + def test_execute_over_an_empty_input_skips(self): + """An empty input skips the task, as ``.expand()`` over nothing does, rather than returning nothing.""" + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator(dag, ListOfDictsExpandInput([]), task_id="empty_input") + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowSkipException, match="empty"): + iterable_op.execute(context=context) + + def test_execute_skip_next_to_a_failure_raises_only_the_failure(self): + """A skipped sub-task is not a failure, so the group holds what failed and nothing else.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [ + {"raise_exception": AirflowSkipException("skip 1")}, + {"raise_exception": RuntimeError("boom")}, + ] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="partial_skip") + + with mock_context(task=iterable_op) as context: + with pytest.raises(RuntimeError, match="boom") as raised: + iterable_op.execute(context=context) + + # A skipped sub-task is not a failure: the only failure is raised on its own. + assert raised.value.__cause__ is None + + def test_execute_skip_next_to_successes_succeeds(self): + """ + A sub-task that skips does not fail the task, as a skipped mapped task instance would not: + the others' results are pushed, the skipped index is left out, and the run counts as complete. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [{"raise_exception": AirflowSkipException("nothing to do")}, {"arg1": 2}] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="skip_and_success") + + with mock_context(task=iterable_op) as context: + materialized = list(iterable_op.execute(context=context)) + store = context["task_state_store"] + + assert materialized == [(2, None, None)] + assert store["_iterable_0"]["status"] == "skipped" + assert store["_iterable_1"]["status"] == "success" + assert store["_iterable_completed"]["completed"] is True + + def test_an_iteration_returning_none_keeps_its_position_and_reads_as_none(self): + """ + Only a skip takes an index out of the sequence. An iteration that returned ``None`` pushed + nothing, as a mapped task instance would, but is counted and reads as ``None``; ``.expand()`` + would not count it, which the docs page says. + """ + + class ValueOrNothing(BaseOperator): + def __init__(self, value=None, skip=False, **kwargs): + super().__init__(**kwargs) + self.value = value + self.skip = skip + + def execute(self, context): + if self.skip: + raise AirflowSkipException("nothing to do") + return self.value + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [{"value": 1}, {"value": None}, {"skip": True}, {"value": 4}] + ) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="none_and_skip", operator_class=ValueOrNothing + ) + + with mock_context(task=iterable_op) as context: + result = iterable_op.execute(context=context) + store = context["task_state_store"] + + assert result.skipped == [2] + assert len(result) == 3 + assert list(result) == [1, None, 4] + assert store["_iterable_1"]["status"] == "success" + assert "result" not in store["_iterable_1"] + + def test_skipped_iteration_is_left_out_of_the_result(self): + """As a skipped mapped task instance, a skipped iteration is not among the values downstream reads.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [{"arg1": 1}, {"raise_exception": AirflowSkipException("nothing to do")}, {"arg1": 3}] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="skip_in_the_middle") + + with mock_context(task=iterable_op) as context: + result = iterable_op.execute(context=context) + + assert result.skipped == [1] + assert len(result) == 2 + assert [value[0] for value in result] == [1, 3] + + @pytest.mark.parametrize( + ("trigger_rule", "skipped"), + [ + ("all_success", True), + ("none_skipped", True), + ("all_done_min_one_success", True), + ("none_failed", False), + ("all_done", False), + ("one_success", False), + ], + ) + def test_partial_skip_skips_the_downstream_tasks_a_skipped_mapped_instance_would( + self, trigger_rule, skipped + ): + """ + A downstream task whose trigger rule skips it when an upstream task instance skipped is skipped, + after the result is pushed for the others; any other downstream task is left to run. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [{"arg1": 1}, {"raise_exception": AirflowSkipException("nothing to do")}] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="partial_skip") + iterable_op >> BaseOperator(task_id="downstream", trigger_rule=trigger_rule) + + with ( + mock_context(task=iterable_op) as context, + patch("airflow.sdk.definitions.iterableoperator._push_xcom_if_needed", autospec=True) as push, + ): + if skipped: + with pytest.raises(DownstreamTasksSkipped) as raised: + iterable_op.execute(context=context) + assert raised.value.tasks == ["downstream"] + (pushed, ti, _), _ = push.call_args + assert ti is context["ti"] + assert [value[0] for value in pushed] == [1] + else: + assert len(iterable_op.execute(context=context)) == 1 + push.assert_not_called() + + def test_no_downstream_task_is_skipped_without_a_skipped_iteration(self): + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]), task_id="no_skip" + ) + iterable_op >> BaseOperator(task_id="downstream") + + with mock_context(task=iterable_op) as context: + assert len(iterable_op.execute(context=context)) == 2 + + def test_clear_after_every_iteration_skipped_runs_every_iteration_again(self): + """ + Regression test: a task whose iterations all skipped is ``SKIPPED``, a final state. Clearing it + must run the iterations again, as clearing a skipped mapped task instance does, instead of + replaying the skips recorded before the clear. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [ + {"raise_exception": AirflowSkipException("nothing to do yet")}, + {"raise_exception": AirflowSkipException("nothing to do yet")}, + ] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="all_skipped_then_cleared") + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + with pytest.raises(AirflowSkipException): + iterable_op.execute(context=context) + + # Cleared: the next attempt runs on the same try_number sequence, now with work to do. + context["ti"].try_number = 2 + with patch.object(MockOperator, "execute", autospec=True, return_value="done") as execute: + result = iterable_op.execute(context=context) + + assert execute.call_count == 2 + assert len(result) == 2 + + @pytest.mark.asyncio + async def test_run_task_does_not_rerun_a_sub_task_skipped_on_a_previous_attempt(self): + """On a retry a ``SKIPPED`` checkpoint is honoured like a ``SUCCESS`` one: the sub-task stays skipped.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="skipped_before") + + with mock_context(task=iterable_op) as context: + task = iterable_op._create_task( + context=context, + index=0, + mapped_kwargs={"raise_exception": RuntimeError("must not run again")}, + jinja_env=iterable_op.get_template_env(dag=dag), + ) + await context["task_state_store"].aset( + task.state_key, + IndexedTaskState( + status=TaskInstanceState.SKIPPED, fingerprint=task.input_fingerprint + ).serialize(), + ) + + with AsyncAwareExecutor(loop=asyncio.get_running_loop(), max_workers=1) as executor: + _, result, raised = await iterable_op._run_task( + executor, + context, + task, + trust_checkpoints=True, + outcomes=IndexedTaskOutcomes(iterable_op, IterationState(), context), + ) + + assert result is None + assert isinstance(raised, AirflowSkipException) + + def test_failed_iteration_checkpoint_records_its_input_and_attempt(self): + """ + A failed iteration's checkpoint carries the input digest and the attempt, like a successful + one, so the next attempt does not take a plain retry for a clear or a changed input. + """ + with DAG("test_dag") as dag: + failing = {"arg1": 1, "fail_on_first_attempt": True} # serializable, so it has a digest + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([failing]), task_id="failed_checkpoint" + ) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 3 + with pytest.raises(RuntimeError): + iterable_op.execute(context=context) + + checkpoint = context["task_state_store"]["_iterable_0"] + + assert checkpoint["status"] == "up_for_retry" + assert IterableOperator._fingerprint(failing) is not None + assert checkpoint["fingerprint"] == IterableOperator._fingerprint(failing) + assert checkpoint["try_number"] == 3 + + def test_parent_timeout_runs_the_retry_callback_while_retries_are_left(self): + """The task is retried for the timeout, so its iteration reports a retry, not a failure.""" + FIRED_CALLBACKS.clear() + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{}]), + task_id="timed_out", + retries=3, + operator_class=MockCallbackTimeoutOperator, + ) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + context["ti"].max_tries = 3 + with pytest.raises(AirflowTaskTimeout): + iterable_op.execute(context=context) + + assert FIRED_CALLBACKS == ["retry"] + + def test_a_sync_items_own_timeout_is_its_failure(self): + """ + A sync operator may raise ``AirflowTaskTimeout`` itself, as a hook that gave up waiting does. + In its worker thread the parent's limit never strikes, so that is the item's own failure, as + it is the mapped task instance's under ``.expand()``: the sibling runs, nothing is killed, and + the task fails with it as with any failed item. + """ + killed: list = [] + + class GivingUpOperator(MockOperator): + def on_kill(self): + killed.append(self.arg1) + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput( + [{"arg1": 1, "raise_exception": AirflowTaskTimeout("gave up waiting")}, {"arg1": 2}] + ), + task_id="own_timeout", + task_concurrency=1, + operator_class=GivingUpOperator, + ) + + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + with pytest.raises(AirflowTaskTimeout, match="gave up waiting"): + iterable_op.execute(context=context) + + assert store["_iterable_0"]["status"] == "up_for_retry" + assert store["_iterable_1"]["status"] == "success" + assert killed == [] + + def test_iteration_cancelled_by_a_sibling_runs_no_callback(self): + """One iteration stops the task; the sibling cancelled on the way did not fail.""" + FIRED_CALLBACKS.clear() + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "defer"}, {"arg1": "sleeper"}]), + task_id="cancelled_sibling", + retries=3, + task_concurrency=2, + operator_class=MockCallbackAsyncOperator, + ) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + context["ti"].max_tries = 3 + with pytest.raises(AirflowFailException, match="attempted to defer"): + iterable_op.execute(context=context) + + assert ("failure", "sleeper") not in FIRED_CALLBACKS + assert ("retry", "sleeper") not in FIRED_CALLBACKS + + def test_failed_publish_keeps_the_success_checkpoint_and_the_retry_only_republishes(self): + """ + When pushing an item's result fails after its SUCCESS checkpoint was written, the checkpoint + stays, and the retry replays the result instead of running the operator again, which would + repeat whatever it did outside Airflow. + """ + runs = [] + original_execute = MockOperator.execute + + def counting_execute(self, context): + runs.append(self.arg1) + return original_execute(self, context) + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{"arg1": 1}]), task_id="failed_publish" + ) + + with ( + mock_context(task=iterable_op) as context, + patch.object(MockOperator, "execute", counting_execute), + ): + context["ti"].try_number = 1 + with patch.object( + IterableOperator, "axcom_push", side_effect=RuntimeError("xcom backend down") + ): + with pytest.raises(RuntimeError, match="xcom backend down"): + iterable_op.execute(context=context) + checkpoint_after_failure = context["task_state_store"]["_iterable_0"]["status"] + + context["ti"].try_number = 2 + result = iterable_op.execute(context=context) + pushed = list(result) + + assert checkpoint_after_failure == "success" + assert runs == [1] + assert pushed == [(1, None, None)] + + def test_the_success_callback_waits_for_the_checkpoint(self): + """ + An item's success callback fires once its SUCCESS checkpoint is written, not when execute returns. + + With the checkpoint write failing, the item announces nothing: the attempt fails, the retry + runs the item again and reports it then, once. A plain task whose result cannot be pushed + fires no success callback either. + """ + CALLBACKS.clear() + runs = [] + original_execute = MockCallbackOperator.execute + + def counting_execute(self, context): + runs.append(self.arg1) + return original_execute(self, context) + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "a"}]), + task_id="checkpoint_fails", + retries=2, + operator_class=MockCallbackOperator, + ) + + with ( + mock_context(task=iterable_op) as context, + patch.object(MockCallbackOperator, "execute", counting_execute), + ): + store = context["task_state_store"] + original_aset = store.aset + failed: list[str] = [] + + async def aset_failing_the_first_success(key, value, **kwargs): + if value.get("status") == "success" and not failed: + failed.append(key) + raise RuntimeError("state store down") + await original_aset(key, value, **kwargs) + + store.aset = aset_failing_the_first_success + context["ti"].try_number = 1 + with pytest.raises(RuntimeError, match="state store down"): + iterable_op.execute(context=context) + fired_after_the_failed_write = list(CALLBACKS) + + context["ti"].try_number = 2 + iterable_op.execute(context=context) + + assert fired_after_the_failed_write == [] + assert runs == ["a", "a"] + assert CALLBACKS == [("success", "a")] + + def test_a_failed_publish_fires_the_success_callback_once(self): + """ + The checkpoint is written before the result is pushed, so a push that fails leaves work the + retry replays rather than runs again: the success callback fired with the checkpoint and + does not fire again on the replay. + """ + CALLBACKS.clear() + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "a"}]), + task_id="publish_fails", + retries=2, + operator_class=MockCallbackOperator, + ) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + with patch.object( + IterableOperator, "axcom_push", side_effect=RuntimeError("xcom backend down") + ): + with pytest.raises(RuntimeError, match="xcom backend down"): + iterable_op.execute(context=context) + fired_after_the_failed_push = list(CALLBACKS) + + context["ti"].try_number = 2 + result = iterable_op.execute(context=context) + pushed = list(result) + + assert fired_after_the_failed_push == [("success", "a")] + assert CALLBACKS == [("success", "a")] + assert pushed == ["a"] + + def test_a_success_callback_that_raises_keeps_the_checkpoint(self): + """ + The success callback runs after the checkpoint is written and may raise what nothing catches: + a ``DeadlockImminentError`` from a synchronous SDK call in an async item's callback. The item's + work is done, so its SUCCESS checkpoint stays and the retry replays the result instead of + running the item again; the callback fired once. + """ + runs: list[str] = [] + fired: list[str] = [] + + class CallbackRaisesOperator(BaseAsyncOperator): + def __init__(self, arg1=None, **kwargs): + kwargs["on_success_callback"] = self._callback + super().__init__(**kwargs) + self.arg1 = arg1 + + def _callback(self, context): + fired.append(self.arg1) + raise DeadlockImminentError("Variable.get on the event loop thread") + + async def aexecute(self, context): + runs.append(self.arg1) + return self.arg1 + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "a"}]), + task_id="callback_raises", + retries=2, + operator_class=CallbackRaisesOperator, + ) + + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + context["ti"].try_number = 1 + with pytest.raises(AirflowFailException, match="synchronous SDK call"): + iterable_op.execute(context=context) + status_after_the_failed_callback = store["_iterable_0"]["status"] + + context["ti"].try_number = 2 + result = iterable_op.execute(context=context) + pushed = list(result) + + assert status_after_the_failed_callback == "success" + assert runs == ["a"] + assert fired == ["a"] + assert pushed == ["a"] + + def test_extra_xcoms_are_checkpointed_once_and_pushed_again_when_a_retry_skips_the_item(self): + """ + The runner deletes every XCom before a retry. An item skipped because it already succeeded + gets its other pushed keys back from its checkpoint, written once with its SUCCESS state. + """ + MockPushingOperator.executed = [] + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "a"}, {"arg1": "b", "fail": True}]), + task_id="extra_xcoms", + task_concurrency=1, + operator_class=MockPushingOperator, + ) + + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + context["ti"].try_number = 1 + with pytest.raises(RuntimeError, match="sibling failed"): + iterable_op.execute(context=context) + checkpoint = store["_iterable_0"] + + # The retry: the failing item succeeds now; the finished one must not run again. + pushed = [] + + async def recording_aset(cls, key, value, **kwargs): + pushed.append((key, value)) + + iterable_op.expand_input = ListOfDictsExpandInput([{"arg1": "a"}, {"arg1": "b"}]) + context["ti"].try_number = 2 + with patch.object(XCom, "aset", classmethod(recording_aset)): + iterable_op.execute(context=context) + + assert checkpoint["status"] == "success" + assert checkpoint["xcoms"] == {"foo": "foo-of-a"} + assert "return_value" not in checkpoint["xcoms"] + assert MockPushingOperator.executed == ["a", "b", "b"] + assert ("foo_0", "foo-of-a") in pushed + assert ("return_value_0", "a") in pushed + + def test_an_item_pushing_many_keys_writes_one_checkpoint(self): + """The extra pushes are kept in memory and written with the one SUCCESS checkpoint, not per push.""" + + class ManyKeysOperator(BaseOperator): + template_fields = () + + def execute(self, context): + for key in range(5): + context["ti"].xcom_push(key=f"key{key}", value=key) + return "done" + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{}]), task_id="many_keys", operator_class=ManyKeysOperator + ) + + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + writes = [] + original_aset = store.aset + + async def counting_aset(key, value, **kwargs): + writes.append(key) + await original_aset(key, value, **kwargs) + + store.aset = counting_aset + iterable_op.execute(context=context) + + assert writes == ["_iterable_0"] + assert store["_iterable_0"]["xcoms"] == {f"key{key}": key for key in range(5)} + + def test_an_async_items_extra_xcoms_are_pushed_again_on_retry(self): + """The async path pushes through axcom_push; its keys are recorded and replayed the same way.""" + + class AsyncPushingOperator(BaseAsyncOperator): + template_fields = ("arg1",) + executed: list = [] + + def __init__(self, arg1=None, fail=False, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + self.fail = fail + + async def aexecute(self, context): + type(self).executed.append(self.arg1) + await context["ti"].axcom_push(key="bar", value=f"bar-of-{self.arg1}") + if self.fail: + raise RuntimeError("sibling failed") + return self.arg1 + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "a"}, {"arg1": "b", "fail": True}]), + task_id="async_extra_xcoms", + task_concurrency=1, + operator_class=AsyncPushingOperator, + ) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + with pytest.raises(RuntimeError, match="sibling failed"): + iterable_op.execute(context=context) + + pushed = [] + + async def recording_aset(cls, key, value, **kwargs): + pushed.append((key, value)) + + iterable_op.expand_input = ListOfDictsExpandInput([{"arg1": "a"}, {"arg1": "b"}]) + context["ti"].try_number = 2 + with patch.object(XCom, "aset", classmethod(recording_aset)): + iterable_op.execute(context=context) + + assert AsyncPushingOperator.executed == ["a", "b", "b"] + assert ("bar_0", "bar-of-a") in pushed + + def test_execute_failed_attempt_leaves_no_completion_marker_so_retry_resumes(self): + """ + Regression test: an attempt that fails must not write the completion marker. + + The marker tells a retry apart from a rerun after a manual clear: when it is present, every + checkpoint is ignored and every index runs again. Raising the collected sub-task failures + outside the ``Checkpoints`` block would let the block exit cleanly and write the marker on a + failed attempt, so the retry would re-run the succeeded indices too instead of resuming. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"raise_exception": ValueError("boom")}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="exec_failed_no_marker") + + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + + with pytest.raises(ValueError, match="boom"): + iterable_op.execute(context=context) + + assert "_iterable_completed" not in store + assert store["_iterable_0"]["status"] == "success" + assert store["_iterable_1"]["status"] == "up_for_retry" + + # The retry: the failing item succeeds now; index 0 must be replayed, not re-run. + context["ti"].try_number = 2 + iterable_op.expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + with patch.object( + IndexedTaskRunner, "run", autospec=True, side_effect=IndexedTaskRunner.run + ) as run: + materialized = list(iterable_op.execute(context=context)) + + assert materialized == [(1, None, None), (2, None, None)] + assert run.call_count == 1 + assert run.call_args.args[0].task_index == 1 # only the failed index runs again + assert store["_iterable_completed"] == {"completed": True, "try_number": 2} + + def test_execute_fail_exception_re_raised_directly_without_retry(self): + """A sub-task that raises AirflowFailException must be re-raised directly (not wrapped in a + BaseExceptionGroup) so the IterableOperator fails without being retried.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [ + {"arg1": 1}, + {"raise_exception": AirflowFailException("boom")}, + ] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="fail_exception", retries=3) + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowFailException, match="boom"): + iterable_op.execute(context=context) + + @pytest.mark.parametrize( + "raised_exception", + [ + DagRunTriggerException( + trigger_dag_id="triggered_dag", + dag_run_id="triggered_run", + conf={}, + reset_dag_run=False, + skip_when_already_exists=False, + wait_for_completion=False, + allowed_states=["success"], + failed_states=["failed"], + poke_interval=1, + deferrable=False, + ), + DownstreamTasksSkipped(tasks=["downstream_task"]), + ], + ids=["DagRunTriggerException", "DownstreamTasksSkipped"], + ) + def test_execute_rejects_trigger_and_downstream_skip_exceptions(self, raised_exception): + """TriggerDagRunOperator (DagRunTriggerException) and downstream-skip operators like + ShortCircuitOperator (DownstreamTasksSkipped) are not supported inside IterableOperator: a + sub-task index has no DAG run or downstream tasks of its own for the trigger/skip to apply + to, so this must fail the whole IterableOperator immediately with a clear error.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"raise_exception": raised_exception}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="unsupported_exception") + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowFailException, match="not supported inside IterableOperator"): + iterable_op.execute(context=context) + + def test_execute_rejects_reschedule_exception(self): + """A reschedule-mode sensor (AirflowRescheduleException) is not supported inside + IterableOperator: the sub-task's index has no task instance of its own to reschedule, so this + must fail the whole IterableOperator immediately with a clear error rather than being silently + aggregated into a retryable BaseExceptionGroup.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [{"raise_exception": AirflowRescheduleException(reschedule_date=None)}] + ) + iterable_op = create_iterable_operator(dag, expand_input, task_id="reschedule_exception") + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowFailException, match="not supported inside IterableOperator"): + iterable_op.execute(context=context) + + @pytest.mark.asyncio + async def test_run_task_skips_sub_task_already_checkpointed_as_succeeded(self): + """ + When the task_state_store already records a sub-task index as succeeded (e.g. because + Airflow retried the whole IterableOperator after a previous partial failure), ``_run_task`` + must skip re-executing that sub-task entirely. + """ + from unittest import mock + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1, "arg2": 10}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="checkpoint_skip") + + with mock_context(task=iterable_op) as context: + jinja_env = iterable_op.get_template_env(dag=dag) + task = iterable_op._create_task( + context=context, index=0, mapped_kwargs={"arg1": 1, "arg2": 10}, jinja_env=jinja_env + ) + task.try_number = 2 # checkpoint is only consulted from the second attempt onwards + await context["task_state_store"].aset( + task.state_key, + IndexedTaskState( + status=TaskInstanceState.SUCCESS, fingerprint=task.input_fingerprint + ).serialize(), + ) + + executor = mock.MagicMock() + _, result, raised = await iterable_op._run_task( + executor, + context, + task, + trust_checkpoints=True, + outcomes=IndexedTaskOutcomes(iterable_op, IterationState(), context), + ) + + assert result is None + assert raised is None + executor.run_sync.assert_not_called() + + @pytest.mark.asyncio + async def test_run_task_reruns_and_checkpoints_success_after_up_for_retry_state(self): + """ + A sub-task whose checkpoint records ``UP_FOR_RETRY`` (left behind by a previous failed or + crashed attempt) is re-executed rather than skipped — only a ``SUCCESS`` checkpoint causes + ``_run_task`` to skip re-execution — and a new ``SUCCESS`` checkpoint recording its result is + stored once it completes. + """ + from airflow.sdk.execution_time.executor import AsyncAwareExecutor + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1, "arg2": 10}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="checkpoint_pending") + + with mock_context(task=iterable_op) as context: + jinja_env = iterable_op.get_template_env(dag=dag) + task = iterable_op._create_task( + context=context, index=0, mapped_kwargs={"arg1": 1, "arg2": 10}, jinja_env=jinja_env + ) + task.try_number = 2 # checkpoint is only consulted from the second attempt onwards + await context["task_state_store"].aset( + task.state_key, + IndexedTaskState(status=TaskInstanceState.UP_FOR_RETRY).serialize(), + ) + + with AsyncAwareExecutor(loop=asyncio.get_running_loop(), max_workers=1) as executor: + result_task, result, raised = await iterable_op._run_task( + executor, + context, + task, + trust_checkpoints=True, + outcomes=IndexedTaskOutcomes(iterable_op, IterationState(), context), + ) + + assert raised is None + assert result == (1, 10, None) + assert ( + result_task.try_number == 2 + ) # try_number is inherited from the parent TI, never mutated + store = context["task_state_store"] + assert ( + store[task.state_key] + == IndexedTaskState( + status=TaskInstanceState.SUCCESS, + result=result, + fingerprint=task.input_fingerprint, + try_number=2, + ).serialize() + ) + + @pytest.mark.asyncio + async def test_run_task_merges_outlet_events_into_shared_context_on_success(self): + """ + A succeeded item's outlet asset events are merged into the task's ``context["outlet_events"]``. + + The task instance reports its outlet events once, from that accessor, when the whole + iteration finishes, so an item's events only reach the server through it. + """ + from airflow.sdk.execution_time.executor import AsyncAwareExecutor + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="outlet_merge", operator_class=MockOutletEventOperator + ) + + with mock_context(task=iterable_op) as context: + jinja_env = iterable_op.get_template_env(dag=dag) + task = iterable_op._create_task( + context=context, index=0, mapped_kwargs={}, jinja_env=jinja_env + ) + + with AsyncAwareExecutor(loop=asyncio.get_running_loop(), max_workers=1) as executor: + _, result, raised = await iterable_op._run_task( + executor, + context, + task, + trust_checkpoints=True, + outcomes=IndexedTaskOutcomes(iterable_op, IterationState(), context), + ) + + assert raised is None + assert result == "done" + accessor = context["outlet_events"][Asset(name="a", uri="s3://bucket/a")] + assert accessor.extra == {"value": "v"} + store = context["task_state_store"] + checkpoint = IndexedTaskState.deserialize(store[task.state_key]) + assert checkpoint.outlet_events == [ + { + "kind": "asset", + "name": "a", + "uri": "s3://bucket/a", + "extra": {"value": "v"}, + "partition_keys": [], + } + ] + + @pytest.mark.asyncio + async def test_run_task_replays_outlet_events_when_skipping_already_succeeded_sub_task(self): + """A sub-task skipped on retry (because its checkpoint already records SUCCESS) never + re-executes, so it would otherwise never re-populate the fresh ``outlet_events`` accessor + created for the new attempt; ``_run_task`` must replay the events it recorded on its + earlier successful attempt instead of silently losing them.""" + from unittest import mock + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="outlet_replay", operator_class=MockOutletEventOperator + ) + + with mock_context(task=iterable_op) as context: + jinja_env = iterable_op.get_template_env(dag=dag) + task = iterable_op._create_task( + context=context, index=0, mapped_kwargs={}, jinja_env=jinja_env + ) + task.try_number = 2 # checkpoint is only consulted from the second attempt onwards + await context["task_state_store"].aset( + task.state_key, + IndexedTaskState( + status=TaskInstanceState.SUCCESS, + fingerprint=task.input_fingerprint, + outlet_events=[ + { + "kind": "asset", + "name": "a", + "uri": "s3://bucket/a", + "extra": {"value": "v"}, + "partition_keys": [], + } + ], + ).serialize(), + ) + + executor = mock.MagicMock() + _, result, raised = await iterable_op._run_task( + executor, + context, + task, + trust_checkpoints=True, + outcomes=IndexedTaskOutcomes(iterable_op, IterationState(), context), + ) + + assert raised is None + assert result is None + executor.run_sync.assert_not_called() + accessor = context["outlet_events"][Asset(name="a", uri="s3://bucket/a")] + assert accessor.extra == {"value": "v"} + + def test_on_kill_propagates_to_active_sub_operators(self): + """IterableOperator.on_kill() (SIGTERM or execution_timeout) must propagate the kill + signal to every sub-task currently in flight, since the default BaseOperator.on_kill() + no-op would otherwise leave running sub-tasks completely unaware of the kill.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="on_kill_test", operator_class=MockOnKillOperator + ) + + active_operator = MockOnKillOperator(task_id="active_sub_task") + iterable_op._state.register(active_operator) + + iterable_op.on_kill() + + assert wait_until(lambda: active_operator.killed) + + def test_on_kill_reaches_sub_operators_that_compare_equal_and_kills_each_once(self): + """ + The sub-operators of one iterated task compare equal (same task_id), so the register is keyed + by identity; and the runner's second on_kill() after a timeout does not kill them again. + """ + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{}]), task_id="on_kill_equal", operator_class=MockOnKillOperator + ) + + first, second = MockOnKillOperator(task_id="same"), MockOnKillOperator(task_id="same") + assert first == second + kills = [] + first.on_kill = lambda: kills.append("first") # type: ignore[method-assign] + second.on_kill = lambda: kills.append("second") # type: ignore[method-assign] + iterable_op._state.register(first) + iterable_op._state.register(second) + + iterable_op.on_kill() + iterable_op.on_kill() + + assert wait_until(lambda: len(kills) == 2) + time.sleep(0.05) # a second kill would arrive here + assert sorted(kills) == ["first", "second"] + + def test_on_kill_reaches_every_sub_operator_when_one_raises_a_base_exception(self): + """A sub-operator's on_kill may raise DeadlockImminentError, a BaseException; the rest are still killed.""" + + class RaisingOnKill(MockOnKillOperator): + def on_kill(self): + raise DeadlockImminentError("sync SDK call on the loop thread") + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{}]), task_id="on_kill_raises", operator_class=MockOnKillOperator + ) + first, second = RaisingOnKill(task_id="first"), MockOnKillOperator(task_id="second") + iterable_op._state.register(first) + iterable_op._state.register(second) + + iterable_op.on_kill() + + assert wait_until(lambda: second.killed) + + def test_on_kill_with_the_loop_paused_kills_off_the_main_thread(self): + """ + Between two ``run_until_complete`` calls the loop is paused and a parked ``asend`` may hold + the comms lock: a sub-operator's sync SDK call in ``on_kill`` on the main thread would wait + for it forever, since only the paused loop can release it. The kill runs in a thread. + """ + seen: dict[str, bool] = {} + done = threading.Event() + main = threading.get_ident() + + class Op(MockOnKillOperator): + def on_kill(self): + seen["on_main"] = threading.get_ident() == main + done.set() + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{}]), task_id="on_kill_paused", operator_class=Op + ) + op = Op(task_id="active") + iterable_op._state.register(op) + + iterable_op.on_kill() # no loop running on this thread, as when the loop is paused + + assert done.wait(5) + assert seen == {"on_main": False} + assert iterable_op._state.stop_requested() + + def test_on_kill_inside_a_running_loop_kills_off_the_loop_thread(self): + """ + The runner's SIGTERM handler calls on_kill on the main thread, usually while the loop runs + there; the sub-operators' on_kill may make a sync SDK call, which must not run on the loop. + """ + seen: dict[str, bool] = {} + done = threading.Event() + main = threading.get_ident() + + class Op(MockOnKillOperator): + def on_kill(self): + seen["on_running_loop"] = asyncio._get_running_loop() is not None + seen["on_main"] = threading.get_ident() == main + done.set() + + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{}]), task_id="on_kill_in_loop", operator_class=Op + ) + op = Op(task_id="active") + iterable_op._state.register(op) + + async def kill_from_the_loop(): + iterable_op.on_kill() + await asyncio.wait_for(asyncio.to_thread(done.wait, 5), 6) + + with event_loop() as loop: + loop.run_until_complete(kill_from_the_loop()) + + assert done.is_set() + assert seen == {"on_running_loop": False, "on_main": False} + + def test_on_kill_takes_no_lock_of_the_iteration_state(self): + """ + The runner's SIGTERM handler runs ``on_kill()`` on the main thread, between two bytecodes of + the loop thread, which may be inside ``register``/``unregister`` and hold the state's lock: + ``on_kill()`` returns at once while the lock is held and the kill follows once it is free. + """ + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{}]), task_id="on_kill_locked", operator_class=MockOnKillOperator + ) + active_operator = MockOnKillOperator(task_id="active_sub_task") + iterable_op._state.register(active_operator) + + with iterable_op._state._lock: + iterable_op.on_kill() + assert iterable_op._state.stop_requested() + assert not active_operator.killed + + assert wait_until(lambda: active_operator.killed) + + def test_on_kill_is_noop_when_no_sub_operators_are_active(self): + """on_kill() must not raise when called with no in-flight sub-tasks (e.g. the + IterableOperator is killed before any sub-task has started, or after all finished).""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="on_kill_noop", operator_class=MockOnKillOperator + ) + + iterable_op.on_kill() # should not raise + + def test_the_kill_thread_is_waited_for_before_the_run_concludes(self): + """ + ``on_kill()`` kills in a thread of its own and returns at once; the run waits for that + thread once the killed indexed tasks came back, so the task does not end, and the process + with it, while a sub-operator's ``on_kill`` is still cleaning up. + """ + executing = threading.Event() + killed = threading.Event() + cleanup_may_finish = threading.Event() + + class SlowToKill(BaseOperator): + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + def execute(self, context): + executing.set() + killed.wait(5) + raise RuntimeError("killed") + + def on_kill(self): + killed.set() + cleanup_may_finish.wait(5) + + outcome: list[BaseException] = [] + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{"arg1": 1}]), task_id="joined_kill", operator_class=SlowToKill + ) + + with mock_context(task=iterable_op) as context: + + def run(): + try: + iterable_op.execute(context=context) + except BaseException as e: + outcome.append(e) + + run_thread = threading.Thread(target=run) + run_thread.start() + assert executing.wait(5) + iterable_op.on_kill() + assert killed.wait(5) + + run_thread.join(0.5) + assert run_thread.is_alive(), "the run concluded while on_kill was still cleaning up" + cleanup_may_finish.set() + run_thread.join(5) + + assert not run_thread.is_alive() + assert isinstance(outcome[0], AirflowTaskTerminated) + + def test_the_wait_for_the_kill_thread_is_bounded(self): + """A sub-operator whose ``on_kill`` never returns holds the run for the shutdown timeout, not forever.""" + executing = threading.Event() + killed = threading.Event() + released = threading.Event() + + class StuckKill(BaseOperator): + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + def execute(self, context): + executing.set() + killed.wait(5) + raise RuntimeError("killed") + + def on_kill(self): + killed.set() + released.wait(30) + + outcome: list[BaseException] = [] + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{"arg1": 1}]), task_id="stuck_kill", operator_class=StuckKill + ) + + with ( + mock_context(task=iterable_op) as context, + patch.object(IterableOperator, "_SHUTDOWN_TIMEOUT", 0.2), + ): + + def run(): + try: + iterable_op.execute(context=context) + except BaseException as e: + outcome.append(e) + + run_thread = threading.Thread(target=run) + run_thread.start() + assert executing.wait(5) + iterable_op.on_kill() + run_thread.join(5) + + released.set() + assert not run_thread.is_alive() + assert isinstance(outcome[0], AirflowTaskTerminated) + + def test_run_task_tracks_active_sub_operator_during_execution(self, monkeypatch: pytest.MonkeyPatch): + """``_run_task`` must register the sub-task's unmapped operator in the iteration state + only for the duration of its execution, so ``on_kill()`` propagates only to sub-tasks that + are actually running.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator( + dag, expand_input, task_id="on_kill_tracking", operator_class=MockOnKillOperator + ) + + with mock_context(task=iterable_op) as context: + jinja_env = iterable_op.get_template_env(dag=dag) + task = iterable_op._create_task( + context=context, index=0, mapped_kwargs={}, jinja_env=jinja_env + ) + + seen_active_during_run = [] + + def tracking_execute(context, ti, log): + seen_active_during_run.append(task.task in iterable_op._state) + + monkeypatch.setattr("airflow.sdk.execution_time.task_runner._execute_task", tracking_execute) + + with event_loop() as loop, AsyncAwareExecutor(loop=loop, max_workers=1) as executor: + _, _, raised = loop.run_until_complete( + iterable_op._run_task( + executor, + context, + task, + trust_checkpoints=False, + outcomes=IndexedTaskOutcomes(iterable_op, IterationState(), context), + ) + ) + + assert raised is None + assert seen_active_during_run == [True] + assert task.task not in iterable_op._state + + def test_multiple_outputs_is_ignored(self): + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}]) + mapped_op = create_mapped_operator(dag=dag, expand_input=expand_input, task_id="multi") + + iterable_op = IterableOperator( + operator=mapped_op, expand_input=expand_input, dag=dag, multiple_outputs=True + ) + + assert iterable_op.multiple_outputs is False + + def test_iterable_execution_timeout_caps_whole_iteration_and_wrapped_operator_retains_it(self): + """IterableOperator keeps execution_timeout as the wall-clock cap on the whole iteration, which + the runner enforces on the outer TI; the wrapped operator keeps the value too, as any of its + own parameters, without a limit of its own per indexed task.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}]) + execution_timeout = timedelta(seconds=7) + mapped_op = create_mapped_operator( + dag, expand_input, task_id="timeout_task", execution_timeout=execution_timeout + ) + + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + assert iterable_op._operator.execution_timeout == execution_timeout + assert iterable_op.execution_timeout == execution_timeout + + @pytest.mark.parametrize( + "base_exception", + [ + SystemExit(1), + KeyboardInterrupt(), + ], + ids=["SystemExit", "KeyboardInterrupt"], + ) + def test_base_exception_not_retried_raises_airflow_fail_exception(self, base_exception): + """ + BaseException subclasses (e.g., SystemExit, KeyboardInterrupt) must never + be retried—they signal conditions where continuing iteration is meaningless. + They should raise AirflowFailException immediately. + """ + with DAG("test_dag") as dag: + # Create a mapped operator that raises a BaseException + expand_input = ListOfDictsExpandInput([{"raise_exception": base_exception}]) + mapped_op = MockOperator.partial( + task_id="base_exception_task", + dag=dag, + retries=3, # Has retries available, but should NOT use them + )._expand( + expand_input, + strict=True, + register_with_dag=False, + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowFailException): + iterable_op.execute(context=context) + + def test_parent_timeout_landing_in_a_sub_task_stays_a_timeout(self): + """ + The parent's ``execution_timeout`` is raised by a signal handler on the main thread, so it + can surface inside whichever async sub-task is running there. It is not that sub-task's + outcome: it has to reach the runner as ``AirflowTaskTimeout``, which retries the task, and + not be turned into a non-retryable ``AirflowFailException`` blamed on the sub-task. + """ + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}]) + iterable_op = create_iterable_operator( + dag, + expand_input, + task_id="timeout_task", + retries=3, + operator_class=MockCallbackTimeoutOperator, + ) + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowTaskTimeout): + iterable_op.execute(context=context) + + assert "_iterable_0" not in context["task_state_store"] + + @pytest.mark.asyncio + async def test_run_task_lets_a_cancellation_through(self): + """A cancelled sub-task stays cancelled: no checkpoint write, no result handed back.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}]) + iterable_op = create_iterable_operator(dag, expand_input, task_id="cancelled_task") + + with mock_context(task=iterable_op) as context: + task = iterable_op._create_task( + context=context, + index=0, + mapped_kwargs={"raise_exception": asyncio.CancelledError()}, + jinja_env=iterable_op.get_template_env(dag=dag), + ) + + with AsyncAwareExecutor(loop=asyncio.get_running_loop(), max_workers=1) as executor: + with pytest.raises(asyncio.CancelledError): + await iterable_op._run_task( + executor, + context, + task, + trust_checkpoints=False, + outcomes=IndexedTaskOutcomes(iterable_op, IterationState(), context), + ) + + assert task.state_key not in context["task_state_store"] + + def test_deadlock_imminent_error_raises_actionable_airflow_fail_exception(self): + """A sub-task that raises DeadlockImminentError (a sync SDK call made from an async + sub-task) must never be retried and must surface an actionable error pointing at async-safe + SDK alternatives, rather than the generic non-Exception BaseException message.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput( + [{"raise_exception": DeadlockImminentError("simulated sync SDK call")}] + ) + mapped_op = MockOperator.partial( + task_id="deadlock_task", + dag=dag, + retries=3, # Has retries available, but should NOT use them + )._expand( + expand_input, + strict=True, + register_with_dag=False, + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowFailException, match="synchronous SDK call"): + iterable_op.execute(context=context) + + def test_deferred_operator_raises_airflow_fail_exception(self): + """A sub-task that raises TaskDeferred must cause IterableOperator to raise AirflowFailException.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}, {}]) + mapped_op = MockDeferredOperator.partial(task_id="deferred_task", dag=dag)._expand( + expand_input, + strict=True, + register_with_dag=False, + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowFailException, match="attempted to defer"): + iterable_op.execute(context=context) + + def test_reschedule_mode_sensor_raises_base_exception_group(self): + """A sub-task that raises AirflowRescheduleException is no longer special-cased: it is treated + like any other sub-task failure and surfaces via BaseExceptionGroup. The requested + reschedule_date is not honored inside IterableOperator — Airflow's standard retry mechanism + (via the IterableOperator's own retries/retry_delay) takes over instead.""" + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}, {}]) + mapped_op = MockRescheduleSensor.partial(task_id="reschedule_sensor", dag=dag)._expand( + expand_input, + strict=True, + register_with_dag=False, + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowFailException, match="attempted to reschedule"): + iterable_op.execute(context=context) + + +class TestIterableOperatorContextIsolation: + """ + Verify that each sub-task run by IterableOperator sees its own indexed + context via get_current_context(), not the parent's. + """ + + def test_subtask_sees_its_own_context(self): + """Each sub-task's get_current_context() must return its own indexed ti, not the parent's.""" + captured: dict[int, object] = {} + + class ContextCapturingOperator(BaseOperator): + def __init__(self, index: int, **kwargs): + super().__init__(**kwargs) + self.index = index + + def execute(self, context): + ctx = get_current_context() + captured[self.index] = ctx["ti"] + return self.index + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"index": 0}, {"index": 1}, {"index": 2}]) + mapped_op = ContextCapturingOperator.partial(task_id="ctx_task", dag=dag)._expand( + expand_input, strict=True, register_with_dag=False + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + iterable_op.execute(context=context) + + parent_ti = context["ti"] + for idx, sub_ti in captured.items(): + # Each sub-task must have seen its own IndexedTaskInstance, not the parent TI. + assert sub_ti is not parent_ti, f"Sub-task {idx} observed the parent context" + assert sub_ti.index == idx, f"Sub-task {idx} observed wrong index {sub_ti.index}" + + +class MockCallbackSyncOperator(BaseOperator): + """Sync twin of MockCallbackAsyncOperator: defers for ``arg1="defer"``, sleeps otherwise.""" + + template_fields = ("arg1",) + + def __init__(self, arg1=None, **kwargs): + kwargs["on_success_callback"] = lambda context: FIRED_CALLBACKS.append(("success", self.arg1)) + kwargs["on_failure_callback"] = lambda context: FIRED_CALLBACKS.append(("failure", self.arg1)) + kwargs["on_retry_callback"] = lambda context: FIRED_CALLBACKS.append(("retry", self.arg1)) + super().__init__(**kwargs) + self.arg1 = arg1 + + def execute(self, context): + if self.arg1 == "defer": + raise TaskDeferred(trigger=None, method_name="execute_complete") # type: ignore[arg-type] + time.sleep(0.5) + return self.arg1 + + +class TestItemThreads: + """ + Which thread an item's code and callbacks run on. + + A sync item runs in a worker thread, enter and exit included, so a sync SDK call from its + ``execute`` or its callbacks waits for the comms lock; on the loop thread the same call raises + ``DeadlockImminentError`` while a sibling's async SDK call is in flight. An async item runs on + the loop and must stay async-safe. + """ + + @staticmethod + def _run(operator_class, is_async): + seen: list[tuple[str, bool, bool]] = [] + main = threading.get_ident() + + def note(where): + seen.append((where, threading.get_ident() != main, asyncio._get_running_loop() is not None)) + + if is_async: + + class Op(operator_class): + def __init__(self, **kwargs): + kwargs["on_success_callback"] = lambda context: note("on_success_callback") + kwargs["on_execute_callback"] = lambda context: note("on_execute_callback") + super().__init__(**kwargs) + + async def aexecute(self, context): + note("aexecute") + + else: + + class Op(operator_class): + def __init__(self, **kwargs): + kwargs["on_success_callback"] = lambda context: note("on_success_callback") + kwargs["on_execute_callback"] = lambda context: note("on_execute_callback") + super().__init__(**kwargs) + + def execute(self, context): + note("execute") + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{}, {}]) + mapped_op = Op.partial(task_id="threads", dag=dag, task_concurrency=2)._expand( + expand_input, strict=True, register_with_dag=False + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + with mock_context(task=iterable_op) as context: + iterable_op.execute(context=context) + assert len(seen) == 6 + return seen + + def test_a_sync_items_code_and_callbacks_run_in_its_worker_thread(self): + for where, off_main, on_running_loop in self._run(BaseOperator, is_async=False): + assert off_main, where + assert not on_running_loop, where + + def test_an_async_items_code_and_callbacks_run_on_the_loop_thread(self): + for where, off_main, on_running_loop in self._run(BaseAsyncOperator, is_async=True): + assert not off_main, where + assert on_running_loop, where + + def test_a_sync_item_cancelled_by_a_sibling_reports_nothing_when_its_thread_finishes(self): + """ + The coroutine waiting for a sync item is cancelled while its thread goes on; the item gets + no checkpoint, so its exit must fire no callback: the next attempt runs it and reports then. + """ + FIRED_CALLBACKS.clear() + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "defer"}, {"arg1": "sleeper"}]), + task_id="cancelled_sync_sibling", + retries=3, + task_concurrency=2, + operator_class=MockCallbackSyncOperator, + ) + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + context["ti"].max_tries = 3 + with pytest.raises(AirflowFailException, match="attempted to defer"): + iterable_op.execute(context=context) + + assert [fired for fired in FIRED_CALLBACKS if fired[1] == "sleeper"] == [] + + +class TestIterationState: + """What one run of an iterated task remembers, apart from the operator's configuration.""" + + def test_registered_operators_are_in_flight_until_unregistered(self): + state = IterationState() + op = MockOperator(task_id="op") + assert op not in state + state.register(op) + assert op in state + state.unregister(op) + assert op not in state + state.unregister(op) # a second time is harmless + + def test_in_flight_is_keyed_by_identity(self): + """Sub-operators of one iterated task compare equal; each is registered on its own.""" + state = IterationState() + first, second = MockOperator(task_id="same"), MockOperator(task_id="same") + assert first == second + state.register(first) + state.register(second) + assert first in state + assert second in state + state.unregister(first) + assert first not in state + assert second in state + + def test_take_in_flight_hands_each_operator_out_once(self): + """The runner's second on_kill() after a timeout must not kill the same operator again.""" + state = IterationState() + first, second, third = (MockOperator(task_id=f"op{i}") for i in range(3)) + state.register(first) + state.register(second) + assert state.take_in_flight() == [first, second] + state.register(third) + assert state.take_in_flight() == [third] + assert state.take_in_flight() == [] + # Still in flight as far as the register goes: a kill does not unregister. + assert first in state + + def test_stop_is_requested_once_asked(self): + state = IterationState() + assert not state.stop_requested() + state.request_stop() + assert state.stop_requested() + state.request_stop() + assert state.stop_requested() + + def test_the_calls_a_signal_handler_makes_take_no_lock(self): + """``request_stop`` and ``start_kill`` run in the SIGTERM handler, on the thread that may hold the lock.""" + state = IterationState() + op = MockOperator(task_id="op") + state.register(op) + taken: list[BaseOperator] = [] + + with state._lock: + state.request_stop() + state.start_kill(lambda: taken.extend(state.take_in_flight())) + assert state.stop_requested() + assert taken == [] + + state.await_kill(5) + assert taken == [op] + + def test_a_copy_is_a_fresh_state(self): + import copy + + state = IterationState() + op = MockOperator(task_id="op") + state.register(op) + state.request_stop() + + copied = copy.deepcopy(state) + + assert copied is not state + assert op not in copied + assert not copied.stop_requested() + assert copied.resolved is None + assert op in state + assert state.stop_requested() + + def test_length_is_unknown_until_the_input_is_resolved(self): + from airflow.sdk.definitions._internal.expandinput import Resolved + + state = IterationState() + assert state.length is None + + async def aget(index): + return {} + + state.resolved = Resolved(3, aget) + assert state.length == 3 + + def test_a_fresh_state_remembers_nothing(self): + state = IterationState() + assert state.resolved is None + assert state.length is None + assert state.take_in_flight() == [] + + +class TestIndexedTaskOutcomes: + """What the indexed tasks ended with, turned into the task's own outcome and its callbacks.""" + + @staticmethod + def _outcomes(retries: int = DEFAULT_RETRIES, state: IterationState | None = None) -> IndexedTaskOutcomes: + with DAG("test_dag") as dag: + op = create_iterable_operator( + dag, ListOfDictsExpandInput([{"arg1": "a"}]), task_id="it", retries=retries + ) + context = {"ti": MagicMock(try_number=1, max_tries=3)} + return IndexedTaskOutcomes(op, state or IterationState(), context) # type: ignore[arg-type] + + @staticmethod + def _resolved(length: int) -> Resolved: + async def aget(index): + return {} + + return Resolved(length, aget) + + @staticmethod + def _task(index: int) -> IndexedTaskInstance: + return MagicMock(spec=IndexedTaskInstance, task_id="it", index=index) + + def test_nothing_raised_is_counted_and_records_nothing_else(self): + outcomes = self._outcomes() + outcomes.record(self._task(0), None) + assert outcomes.total == 1 + assert outcomes.exceptions == [] + assert outcomes.skipped == {} + + def test_a_skip_is_noted_under_its_index(self): + outcomes = self._outcomes() + skip = AirflowSkipException("odd") + outcomes.record(self._task(2), skip) + assert outcomes.skipped == {2: skip} + assert outcomes.exceptions == [] + + def test_a_failure_is_collected(self): + outcomes = self._outcomes() + raised = ValueError("boom") + outcomes.record(self._task(1), raised) + assert outcomes.exceptions == [raised] + + @pytest.mark.parametrize( + ("raised", "match"), + [ + (TaskDeferred(trigger=MagicMock(), method_name="x"), "attempted to defer"), + (AirflowRescheduleException(datetime.now(UTC)), "attempted to reschedule"), + (DownstreamTasksSkipped(tasks=["a"]), "raised DownstreamTasksSkipped"), + (DeadlockImminentError("sync call on the loop"), "made a synchronous SDK call"), + (KeyboardInterrupt(), "raised a non-Exception BaseException"), + ], + ) + def test_an_outcome_the_iteration_cannot_carry_ends_the_task_at_once(self, raised, match): + """Every later indexed task would end the same way, so nothing is collected: it is raised.""" + outcomes = self._outcomes() + with pytest.raises(AirflowFailException, match=rf"Sub-task it\[3\] {match}") as info: + outcomes.record(self._task(3), raised) + assert info.value.__cause__ is raised + assert outcomes.exceptions == [] + + def test_a_kill_concludes_as_terminated_whatever_else_happened(self): + state = IterationState() + state.resolved = self._resolved(5) + outcomes = self._outcomes(state=state) + outcomes.record(self._task(0), ValueError("a")) + outcomes.record(self._task(1), None) + state.request_stop() + with pytest.raises(AirflowTaskTerminated, match="2 of 5 items ran") as info: + outcomes.conclude() + assert isinstance(info.value.__cause__, BaseExceptionGroup) + assert info.value.__cause__.exceptions == (outcomes.exceptions[0],) + + def test_an_item_not_started_after_the_kill_is_counted_apart(self): + """It ran no code: not a failure to collect or log, not an item that ran, named in the message.""" + state = IterationState() + state.resolved = self._resolved(3) + outcomes = self._outcomes(state=state) + outcomes.record(self._task(0), None) + outcomes.record(self._task(1), IndexedTaskInstanceNotStarted("pulled before the kill")) + assert outcomes.total == 1 + assert outcomes.not_started == 1 + assert outcomes.exceptions == [] + state.request_stop() + with pytest.raises( + AirflowTaskTerminated, match=r"1 of 3 items ran, 1 pulled but never started, 1 never pulled\.$" + ): + outcomes.conclude() + + def test_a_kill_after_every_item_was_pulled_claims_no_remainder(self): + """Every item pulled and run, the killed one back as a failure: the message ends with the count.""" + state = IterationState() + state.resolved = self._resolved(2) + outcomes = self._outcomes(state=state) + outcomes.record(self._task(0), ValueError("killed")) + outcomes.record(self._task(1), None) + state.request_stop() + with pytest.raises(AirflowTaskTerminated, match=r"2 of 2 items ran\.$"): + outcomes.conclude() + + def test_a_kill_while_resolving_is_not_an_empty_input(self): + state = IterationState() + state.request_stop() + with pytest.raises(AirflowTaskTerminated, match="before its input was resolved") as info: + self._outcomes(state=state).conclude() + assert info.value.__cause__ is None + + def test_failures_conclude_on_the_exception_handed_to_the_runner(self): + outcomes = self._outcomes() + raised = ValueError("a") + outcomes.record(self._task(0), raised) + outcomes.record(self._task(1), None) + with pytest.raises(ValueError, match="a") as info: + outcomes.conclude() + assert info.value is raised + + def test_an_empty_input_concludes_as_a_skip(self): + with pytest.raises(AirflowSkipException, match="empty"): + self._outcomes().conclude() + + def test_every_indexed_task_skipped_skips_the_task(self): + outcomes = self._outcomes() + first = AirflowSkipException("first") + outcomes.record(self._task(0), first) + outcomes.record(self._task(1), AirflowSkipException("second")) + with pytest.raises(AirflowSkipException) as info: + outcomes.conclude() + assert info.value is first + + def test_a_partial_skip_concludes_with_the_skipped_indices_in_order(self): + outcomes = self._outcomes() + outcomes.record(self._task(2), AirflowSkipException("c")) + outcomes.record(self._task(0), AirflowSkipException("a")) + outcomes.record(self._task(1), None) + assert outcomes.conclude() == [0, 2] + + def test_failed_runners_are_kept_in_the_order_they_failed(self): + outcomes = self._outcomes() + first, second = object(), object() + outcomes.note_failed(first) # type: ignore[arg-type] + outcomes.note_failed(second) # type: ignore[arg-type] + assert outcomes.failed_runners == (first, second) + + def test_the_failed_runners_report_the_tasks_fate_when_an_exception_leaves_the_block(self): + """A retry while attempts are left, a final failure for a fail-fast exception, in failure order.""" + for raised, will_retry in [(ValueError("x"), True), (AirflowFailException("x"), False)]: + outcomes = self._outcomes(retries=2) + first, second = MagicMock(), MagicMock() + first.task_instance.is_eligible_to_retry = True + outcomes.note_failed(first) + outcomes.note_failed(second) + manager = Mock() + manager.attach_mock(first, "first") + manager.attach_mock(second, "second") + + with pytest.raises(type(raised)), outcomes: + raise raised + + assert manager.mock_calls == [ + call.first.report_failure(task_will_retry=will_retry), + call.second.report_failure(task_will_retry=will_retry), + ] + + def test_nothing_is_reported_when_the_block_ends_normally(self): + """No failure left the block, so no indexed task's callback is waiting on it.""" + outcomes = self._outcomes() + runner = MagicMock() + outcomes.note_failed(runner) + with outcomes: + pass + runner.report_failure.assert_not_called() + + +KILL_TARGET: list = [] + + +class MockKillingOperator(BaseOperator): + """Operator whose ``arg1="kill"`` item kills the iterated task from inside, as SIGTERM would.""" + + template_fields = ("arg1",) + + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + def execute(self, context): + if self.arg1 == "kill": + KILL_TARGET[0].on_kill() + return self.arg1 + + +class MockEchoAsyncOperator(BaseAsyncOperator): + """Async operator that returns its ``arg1``, for kills that come from elsewhere than the item.""" + + template_fields = ("arg1",) + + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + async def aexecute(self, context): + return self.arg1 + + +class TestAKillSticks: + """on_kill() stops the iteration: nothing new starts, and the task fails without a retry.""" + + def test_a_kill_stops_the_iteration_and_fails_the_task(self): + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": "kill"}, {"arg1": 1}, {"arg1": 2}, {"arg1": 3}]), + task_id="killed", + task_concurrency=1, + operator_class=MockKillingOperator, + ) + KILL_TARGET[:] = [iterable_op] + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + with pytest.raises(AirflowTaskTerminated, match="1 of 4 items ran"): + iterable_op.execute(context=context) + + assert "_iterable_completed" not in store + assert store["_iterable_0"]["status"] == "success" + assert all(f"_iterable_{index}" not in store for index in (1, 2, 3)) + + def test_an_item_pulled_before_the_kill_but_not_started_does_not_run(self): + """ + on_kill() reaches what has started; an item whose checkpoint read is in flight has not. + + On a retry attempt every item reads its checkpoint before its runner starts. A kill that + lands during that read (here from the read itself, as a SIGTERM would land on the main + thread) finds the item unregistered: it must not run afterwards, gets no checkpoint and + fires no callback, as the items never pulled do not, and the message counts it apart. + """ + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": 0}, {"arg1": 1}, {"arg1": 2}]), + task_id="killed", + task_concurrency=3, + operator_class=MockEchoAsyncOperator, + ) + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 2 + store = context["task_state_store"] + read = store.aget + + async def aget_and_kill(key, default=None): + if key == "_iterable_1": + iterable_op.on_kill() + return await read(key, default) + + store.aget = aget_and_kill + with pytest.raises( + AirflowTaskTerminated, match="1 of 3 items ran, 1 pulled but never started" + ): + iterable_op.execute(context=context) + + assert "_iterable_completed" not in store + assert store["_iterable_0"]["status"] == "success" + assert all(f"_iterable_{index}" not in store for index in (1, 2)) + + def test_a_kill_before_the_run_started_still_stops_it(self): + """SIGTERM can land between the operator's creation and _run_tasks: the stop must survive.""" + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]), task_id="killed_early" + ) + iterable_op.on_kill() + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + with pytest.raises(AirflowTaskTerminated, match="before its input was resolved"): + iterable_op.execute(context=context) + + assert "_iterable_completed" not in store + assert "_iterable_0" not in store + + def test_a_rerun_in_the_same_process_does_not_see_the_earlier_kill(self): + """dag.test() and the tests run one operator object several times; the state is per run.""" + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]), task_id="killed_then_rerun" + ) + iterable_op.on_kill() + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowTaskTerminated): + iterable_op.execute(context=context) + + result = iterable_op.execute(context=context) + + assert len(result) == 2 + assert "_iterable_completed" in context["task_state_store"] + + def test_a_kill_before_any_item_started_is_not_an_empty_input(self): + """A kill while the input is resolved leaves nothing to run; that is a kill, not a skip.""" + xcom_arg = make_xcom_arg(None) + + async def aresolve(*a, **kw): + KILL_TARGET[0].on_kill() + return [{"arg1": 1}, {"arg1": 2}] + + xcom_arg.aresolve = aresolve + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput(xcom_arg), + task_id="killed_early", + operator_class=MockKillingOperator, + ) + KILL_TARGET[:] = [iterable_op] + with mock_context(task=iterable_op) as context: + store = context["task_state_store"] + with pytest.raises(AirflowTaskTerminated, match="0 of 2 items ran"): + iterable_op.execute(context=context) + + assert "_iterable_completed" not in store + + +KILLED_ON_TIMEOUT: list = [] + + +class MockSlowSyncOperator(BaseOperator): + """Sync operator that waits up to 3 s for on_kill(), as an operator stopping an external job would.""" + + template_fields = ("arg1",) + + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + self.stop = threading.Event() + + def execute(self, context): + self.stop.wait(3) + + def on_kill(self): + KILLED_ON_TIMEOUT.append(("sync", self.arg1)) + self.stop.set() + + +class MockSlowAsyncKillableOperator(BaseAsyncOperator): + template_fields = ("arg1",) + + def __init__(self, arg1=None, **kwargs): + super().__init__(**kwargs) + self.arg1 = arg1 + + async def aexecute(self, context): + await asyncio.sleep(3) + + def on_kill(self): + KILLED_ON_TIMEOUT.append(("async", self.arg1)) + + +class TestExecutionTimeoutKillsInFlightSubTasks: + """ + The parent's execution_timeout reaches every sub-task in flight, once, through on_kill(). + + Runs the operator through the runner's own ``_run_execute_callable`` with a real timeout, as the + task runner does, so the order in which the executor cancels and the runner calls on_kill() is + the real one. + """ + + @pytest.mark.parametrize( + ("operator_class", "kind"), + [(MockSlowSyncOperator, "sync"), (MockSlowAsyncKillableOperator, "async")], + ids=["sync", "async"], + ) + def test_every_sub_task_in_flight_is_killed_once(self, operator_class, kind, mock_supervisor_comms): + from airflow.sdk.execution_time.task_runner import _run_execute_callable + + KILLED_ON_TIMEOUT.clear() + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + mapped_op = create_mapped_operator( + dag, + expand_input, + task_id="timed_out", + task_concurrency=2, + execution_timeout=timedelta(milliseconds=300), + operator_class=operator_class, + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowTaskTimeout): + _run_execute_callable(context, iterable_op.execute, iterable_op) + + assert sorted(KILLED_ON_TIMEOUT) == [(kind, 1), (kind, 2)] + + def test_the_parents_timeout_is_sent_to_the_supervisor_once(self, mock_supervisor_comms): + """ + Sync items with an execution_timeout run in worker threads, where TimeoutPosix cannot fire, + under the parent's limit. Re-sending SetExecutionTimeout per item would move the + supervisor's hard-kill deadline to the last item started, up to one full timeout late. + """ + from airflow.sdk.execution_time.comms import SetExecutionTimeout + from airflow.sdk.execution_time.task_runner import _run_execute_callable + + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}, {"arg1": 3}]) + mapped_op = create_mapped_operator( + dag, + expand_input, + task_id="timed_items", + task_concurrency=2, + execution_timeout=timedelta(seconds=30), + ) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + _run_execute_callable(context, iterable_op.execute, iterable_op) + + sent = [ + call.kwargs.get("msg", call.args[0] if call.args else None) + for call in mock_supervisor_comms.send.call_args_list + ] + assert [msg for msg in sent if isinstance(msg, SetExecutionTimeout)] == [ + SetExecutionTimeout(timeout_seconds=30.0) + ] + + def test_the_struck_async_items_on_kill_runs_off_the_loop_thread(self): + """ + The parent's timeout can land inside the async item running on the loop thread. Its on_kill + may make a sync SDK call (cancelling a remote job through a sync hook), which raises + DeadlockImminentError on that thread; killed there, the timeout would turn into a failure + without a retry and the job would go on. The struck item stays registered and the parent + kills it off the loop with the others. The strike is simulated by the item raising the + timeout itself, as the signal handler would inside its coroutine. + """ + + class SyncHookOnKill(MockSlowAsyncKillableOperator): + async def aexecute(self, context): + if self.arg1 == 1: + raise AirflowTaskTimeout("the parent ran out of time") + await asyncio.sleep(3) + + def on_kill(self): + if asyncio._get_running_loop() is not None: + raise DeadlockImminentError("sync SDK call on the loop thread") + super().on_kill() + + KILLED_ON_TIMEOUT.clear() + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator( + dag, + ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]), + task_id="timed_out_sync_hook", + task_concurrency=2, + operator_class=SyncHookOnKill, + ) + with mock_context(task=iterable_op) as context: + with pytest.raises(AirflowTaskTimeout): + iterable_op.execute(context=context) + + assert sorted(KILLED_ON_TIMEOUT) == [("async", 1), ("async", 2)] + + +class TestFailureHandedToTheRunner: + """ + The runner classifies the outcome by exception type, so an item's own exception reaches it + whenever one decides: a fail-fast one, the only one, or the one the retry policy decides on. + """ + + @staticmethod + def _execute(items, retry_policy=None): + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput(items) + mapped_op = MockOperator.partial( + task_id="failing", dag=dag, retries=2, retry_policy=retry_policy, task_concurrency=1 + )._expand(expand_input, strict=True, register_with_dag=False) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = 1 + context["ti"].max_tries = 2 + try: + iterable_op.execute(context=context) + except BaseException as exc: + return exc + raise AssertionError("execute() did not raise") + + def test_sensor_timeout_is_raised_on_its_own_next_to_other_failures(self): + from airflow.sdk.exceptions import AirflowSensorTimeout + + raised = self._execute( + [{"raise_exception": ValueError("boom")}, {"raise_exception": AirflowSensorTimeout("poked out")}] + ) + + assert isinstance(raised, AirflowSensorTimeout) + assert isinstance(raised.__cause__, BaseExceptionGroup) + assert sorted(type(exc).__name__ for exc in raised.__cause__.exceptions) == [ + "AirflowSensorTimeout", + "ValueError", + ] + + def test_retry_policy_decides_on_the_items_own_exception(self): + from airflow.sdk.definitions.retry_policy import ExceptionRetryPolicy, RetryAction, RetryRule + + policy = ExceptionRetryPolicy(rules=[RetryRule(exception=PermissionError, action=RetryAction.FAIL)]) + + raised = self._execute( + [ + {"arg1": 1}, + {"raise_exception": ValueError("boom")}, + {"raise_exception": PermissionError("no")}, + ], + retry_policy=policy, + ) + + assert isinstance(raised, PermissionError) + assert isinstance(raised.__cause__, BaseExceptionGroup) + # What the runner evaluates next: the rule matches the item's exception, never the group. + assert policy.evaluate(exception=raised, try_number=1, max_tries=2).action == RetryAction.FAIL + assert ( + policy.evaluate(exception=raised.__cause__, try_number=1, max_tries=2).action + == RetryAction.DEFAULT + ) + + def test_several_failures_no_policy_decides_on_stay_a_group(self): + raised = self._execute([{"raise_exception": ValueError("one")}, {"raise_exception": KeyError("two")}]) + + assert isinstance(raised, BaseExceptionGroup) + assert sorted(type(exc).__name__ for exc in raised.exceptions) == ["KeyError", "ValueError"] + + +class TestFingerprintCoversPartialInputsFromUpstream: + """ + A checkpoint is honoured only for the input it was written for, and that input includes the + ``.partial()`` kwargs an upstream task provides: clearing the upstream can change them. + """ + + @staticmethod + def _run_twice(partial_value_on_retry, arg2_first="from-upstream", arg2_is_upstream=True): + """Attempt 1: item 2 fails. Attempt 2 (retry or clear): item 2 succeeds. Returns runs and results.""" + upstream_value = [arg2_first] + runs = [] + original_execute = MockOperator.execute + + def counting_execute(self, context): + runs.append((context["ti"].try_number, self.arg1)) + return original_execute(self, context) + + with DAG("test_dag") as dag: + if arg2_is_upstream: + arg2 = make_xcom_arg(None) + arg2.resolve = lambda *args, **kwargs: upstream_value[0] + else: + arg2 = arg2_first + expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2, "fail_on_first_attempt": True}]) + mapped_op = MockOperator.partial( + task_id="partial_input", dag=dag, arg2=arg2, task_concurrency=1 + )._expand(expand_input, strict=True, register_with_dag=False) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with ( + mock_context(task=iterable_op) as context, + patch.object(MockOperator, "execute", counting_execute), + ): + context["ti"].try_number = 1 + with pytest.raises(RuntimeError): + iterable_op.execute(context=context) + upstream_value[0] = partial_value_on_retry + iterable_op.expand_input = ListOfDictsExpandInput([{"arg1": 1}, {"arg1": 2}]) + context["ti"].try_number = 2 + results = list(iterable_op.execute(context=context)) + return runs, results + + def test_a_changed_upstream_value_runs_the_item_again(self): + runs, results = self._run_twice("changed-upstream") + + assert (2, 1) in runs + assert results == [(1, "changed-upstream", None), (2, "changed-upstream", None)] + + def test_an_unchanged_upstream_value_keeps_the_checkpoint(self): + runs, results = self._run_twice("from-upstream") + + assert (2, 1) not in runs + assert results == [(1, "from-upstream", None), (2, "from-upstream", None)] + + def test_a_templated_partial_value_does_not_make_checkpoints_stale(self): + """Only upstream values count: a value rendered anew every attempt must not force a rerun.""" + runs, results = self._run_twice(None, arg2_first="{{ ti.try_number }}", arg2_is_upstream=False) + + assert (2, 1) not in runs + assert [value[0] for value in results] == [1, 2] + + def test_partial_inputs_from_upstream_are_read_from_the_rendered_operator(self): + upstream = make_xcom_arg(None) + rendered = SimpleNamespace(arg2="rendered-arg2", op_kwargs={"y": "rendered-y", "z": 1}, retries=2) + with DAG("test_dag") as dag: + iterable_op = create_iterable_operator(dag, ListOfDictsExpandInput([{"arg1": "a"}])) + iterable_op.partial_kwargs = {"arg2": upstream, "op_kwargs": {"y": upstream, "z": 1}, "retries": 2} + + inputs = iterable_op._partial_inputs_from_upstream(rendered) + + assert inputs == {"arg2": "rendered-arg2", "op_kwargs.y": "rendered-y"} + + +CALLBACKS: list = [] + + +class MockCallbackOperator(BaseOperator): + """Operator that records which callback each of its items gets.""" + + template_fields = ("arg1",) + arg1: Any + + def __init__(self, arg1=None, raise_exception: BaseException | None = None, **kwargs): + kwargs["on_success_callback"] = lambda context: CALLBACKS.append(("success", self.arg1)) + kwargs["on_failure_callback"] = lambda context: CALLBACKS.append(("failure", self.arg1)) + kwargs["on_retry_callback"] = lambda context: CALLBACKS.append(("retry", self.arg1)) + super().__init__(**kwargs) + self.arg1 = arg1 + self.raise_exception = raise_exception + + def execute(self, context): + if self.raise_exception is not None: + raise self.raise_exception + return self.arg1 + + +class TestCallbacksFollowTheTasksFate: + """ + A failed item's failure or retry callback waits until every item has run, and then says what + happens to the task: retried, or failed for good. Success callbacks fire right away. + """ + + @staticmethod + def _run(items, try_number=1, max_tries=2, retry_policy=None): + CALLBACKS.clear() + with DAG("test_dag") as dag: + expand_input = ListOfDictsExpandInput(items) + mapped_op = MockCallbackOperator.partial( + task_id="callbacks", dag=dag, retries=2, retry_policy=retry_policy, task_concurrency=1 + )._expand(expand_input, strict=True, register_with_dag=False) + iterable_op = IterableOperator(operator=mapped_op, expand_input=expand_input, dag=dag) + + with mock_context(task=iterable_op) as context: + context["ti"].try_number = try_number + context["ti"].max_tries = max_tries + try: + iterable_op.execute(context=context) + except BaseException as exc: + return exc, list(CALLBACKS) + return None, list(CALLBACKS) + + def test_a_siblings_fail_exception_turns_every_failure_into_a_final_one(self): + """ + A failed item's retry callback waits until no sibling rules the retry out. + + A ``ValueError`` item next to an ``AirflowFailException`` item is reported as a failure: + the task fails without a retry, so announcing a retry for the first item would be wrong. + """ + raised, fired = self._run( + [ + {"arg1": "ok"}, + {"arg1": "value_error", "raise_exception": ValueError("x")}, + {"arg1": "fail", "raise_exception": AirflowFailException("stop")}, + ] + ) + + assert isinstance(raised, AirflowFailException) + assert fired == [("success", "ok"), ("failure", "value_error"), ("failure", "fail")] + + def test_failures_the_task_is_retried_for_all_get_the_retry_callback(self): + raised, fired = self._run( + [ + {"arg1": "a", "raise_exception": ValueError("a")}, + {"arg1": "b", "raise_exception": KeyError("b")}, + ] + ) + + assert isinstance(raised, BaseExceptionGroup) + assert fired == [("retry", "a"), ("retry", "b")] + + def test_on_the_last_attempt_every_failure_is_final(self): + _, fired = self._run( + [ + {"arg1": "a", "raise_exception": ValueError("a")}, + {"arg1": "b", "raise_exception": KeyError("b")}, + ], + try_number=3, + max_tries=2, + ) + + assert fired == [("failure", "a"), ("failure", "b")] + + def test_the_retry_policy_is_evaluated_once_per_failed_item(self): + """ + Choosing the exception for the runner costs one evaluation per failed item, and no more. + + The callbacks reuse the decision taken for the exception handed over instead of evaluating + the policy again: a policy that calls a model (common.ai's ``LLMRetryPolicy``) pays per call + and may answer differently each time. + """ + from airflow.sdk.definitions.retry_policy import ExceptionRetryPolicy, RetryDecision + + class CountingPolicy(ExceptionRetryPolicy): + calls: list = [] + + def evaluate(self, exception, try_number, max_tries, context=None): + self.calls.append(exception) + return RetryDecision.retry() + + _, fired = self._run( + [ + {"arg1": "a", "raise_exception": ValueError("a")}, + {"arg1": "b", "raise_exception": KeyError("b")}, + {"arg1": "c", "raise_exception": OSError("c")}, + ], + retry_policy=CountingPolicy(rules=[]), + ) + + assert len(CountingPolicy.calls) == 3 + assert fired == [("retry", "a"), ("retry", "b"), ("retry", "c")] + + def test_an_undecided_group_is_not_evaluated_again_for_the_callbacks(self): + """ + Failures whose decisions are all the default are handed over as a group, decided once. + + Choosing costs one evaluation per failed item; the group then carries the default, so the + callbacks follow it without evaluating the policy once more, on an exception no item raised. + """ + from airflow.sdk.definitions.retry_policy import ExceptionRetryPolicy, RetryDecision + + class DefaultPolicy(ExceptionRetryPolicy): + calls: list = [] + + def evaluate(self, exception, try_number, max_tries, context=None): + self.calls.append(exception) + return RetryDecision.default() + + raised, fired = self._run( + [ + {"arg1": "a", "raise_exception": ValueError("a")}, + {"arg1": "b", "raise_exception": KeyError("b")}, + {"arg1": "c", "raise_exception": OSError("c")}, + ], + retry_policy=DefaultPolicy(rules=[]), + ) + + assert isinstance(raised, BaseExceptionGroup) + assert len(DefaultPolicy.calls) == 3 + assert fired == [("retry", "a"), ("retry", "b"), ("retry", "c")] + + def test_the_callbacks_follow_the_decision_taken_for_the_exception_handed_over(self): + """ + A policy that answers differently per call cannot make the callbacks disagree with the task. + + The second item's exception wins with FAIL and is handed to the runner; evaluating the policy + once more for the callbacks would get RETRY and announce a retry the runner does not take. + """ + from airflow.sdk.definitions.retry_policy import ExceptionRetryPolicy, RetryDecision + + class AlternatingPolicy(ExceptionRetryPolicy): + answers = iter([RetryDecision.retry(), RetryDecision.fail("no"), RetryDecision.retry()]) + + def evaluate(self, exception, try_number, max_tries, context=None): + return next(self.answers) + + raised, fired = self._run( + [ + {"arg1": "a", "raise_exception": ValueError("a")}, + {"arg1": "b", "raise_exception": KeyError("b")}, + ], + retry_policy=AlternatingPolicy(rules=[]), + ) + + assert isinstance(raised, KeyError) + assert fired == [("failure", "a"), ("failure", "b")] + + def test_a_retry_policy_that_fails_the_task_makes_every_failure_final(self): + from airflow.sdk.definitions.retry_policy import ExceptionRetryPolicy, RetryAction, RetryRule + + policy = ExceptionRetryPolicy(rules=[RetryRule(exception=PermissionError, action=RetryAction.FAIL)]) + + _, fired = self._run( + [ + {"arg1": "value_error", "raise_exception": ValueError("x")}, + {"arg1": "denied", "raise_exception": PermissionError("no")}, + ], + retry_policy=policy, + ) + + assert fired == [("failure", "value_error"), ("failure", "denied")] + + def test_an_item_the_operator_rejects_fails_the_task_for_every_item(self): + """A downstream skip from an item is rejected with AirflowFailException, so nothing is retried.""" + raised, fired = self._run( + [ + {"arg1": "value_error", "raise_exception": ValueError("x")}, + {"arg1": "skipper", "raise_exception": DownstreamTasksSkipped(tasks=["downstream"])}, + ] + ) + + assert isinstance(raised, AirflowFailException) + assert fired == [("failure", "value_error"), ("failure", "skipper")] + + def test_a_failed_items_callback_waits_for_the_items_after_it(self): + _, fired = self._run([{"arg1": "first", "raise_exception": ValueError("x")}, {"arg1": "second"}]) + + assert fired == [("success", "second"), ("retry", "first")] + + +class TestIterableOperatorCopy: + """An iterated task can be deep-copied, as dag.partial_subset() does for every task it keeps.""" + + @staticmethod + def _dag(): + from airflow.sdk import task + + with DAG("copy_dag") as dag: + + @task + def up(): + return [1, 2] + + @task + def f(x): + return x + + @task + def down(values): + return values + + down(f.iterate(x=up())) + return dag + + def test_deepcopy_gets_its_own_state_with_no_sub_tasks_in_flight(self): + import copy + + iterable_op = self._dag().task_dict["f"] + in_flight = MockOperator(task_id="in_flight") + iterable_op._state.register(in_flight) + iterable_op._state.request_stop() + + copied = copy.deepcopy(iterable_op) + + assert isinstance(copied, IterableOperator) + assert copied.task_id == "f" + assert copied._state is not iterable_op._state + assert in_flight not in copied._state + assert not copied._state.stop_requested() + # The original is untouched. + assert in_flight in iterable_op._state + assert iterable_op._state.stop_requested() + + def test_the_copy_prepared_for_execution_gets_its_own_state(self): + """ + ``prepare_for_execution`` copies with ``copy.copy``. A kill in one attempt must not stop the + next one, which ``dag.test()`` prepares from the same Dag operator. + """ + iterable_op = self._dag().task_dict["f"] + + first = iterable_op.prepare_for_execution() + first.on_kill() + second = iterable_op.prepare_for_execution() + + assert first._state is not iterable_op._state + assert first._state.stop_requested() + assert not iterable_op._state.stop_requested() + assert second._state is not first._state + assert not second._state.stop_requested() + assert second._lock_for_execution + + def test_partial_subset_keeps_the_iterated_task(self): + dag = self._dag() + + subset = dag.partial_subset("f", include_upstream=True, include_downstream=True) + + assert sorted(subset.task_dict) == ["down", "f", "up"] + assert isinstance(subset.task_dict["f"], IterableOperator) + assert subset.task_dict["f"] is not dag.task_dict["f"] + + +class TestCheckpoints: + @staticmethod + def _context(try_number: int, marker: dict | None = None): + from airflow.sdk.execution_time.context import TaskStateStoreAccessor + + store = create_autospec(TaskStateStoreAccessor, instance=True) + store.get.return_value = marker + ti = SimpleNamespace(task_id="my_task", try_number=try_number) + return {"task_state_store": store, "ti": ti}, store + + def test_first_attempt_never_reads_and_marks_completed_on_exit(self): + context, store = self._context(try_number=1) + + with Checkpoints(context) as checkpoints: + assert checkpoints.trust_checkpoints is False + + store.get.assert_not_called() + store.delete.assert_not_called() + store.set.assert_called_once_with("_iterable_completed", {"completed": True, "try_number": 1}) + + def test_retry_without_marker_trusts_checkpoints(self): + context, store = self._context(try_number=2, marker=None) + + with Checkpoints(context) as checkpoints: + assert checkpoints.trust_checkpoints is True + + store.get.assert_called_once_with("_iterable_completed") + store.delete.assert_not_called() + store.set.assert_called_once_with("_iterable_completed", {"completed": True, "try_number": 2}) + + def test_failed_attempt_leaves_no_marker(self): + context, store = self._context(try_number=1) + + with pytest.raises(RuntimeError, match="sub-task failures"): + with Checkpoints(context): + raise RuntimeError("sub-task failures") + + store.set.assert_not_called() + + def test_attempt_where_every_iteration_skipped_marks_completed(self): + """The task ends ``SKIPPED``; a clear of it must run every iteration again.""" + context, store = self._context(try_number=1) + + with pytest.raises(AirflowSkipException): + with Checkpoints(context): + raise AirflowSkipException("every iteration skipped") + + store.set.assert_called_once_with("_iterable_completed", {"completed": True, "try_number": 1}) + + def test_rerun_after_clear_records_where_it_starts_and_ignores_checkpoints(self): + context, store = self._context(try_number=2, marker={"completed": True, "try_number": 1}) + + with Checkpoints(context) as checkpoints: + assert checkpoints.trust_checkpoints is False + store.set.assert_called_once_with("_iterable_completed", {"completed": False, "since": 2}) + + store.set.assert_called_with("_iterable_completed", {"completed": True, "try_number": 2}) + + def test_failed_rerun_leaves_where_it_started(self): + context, store = self._context(try_number=2, marker={"completed": True, "try_number": 1}) + + with pytest.raises(RuntimeError): + with Checkpoints(context): + raise RuntimeError("sub-task failed") + + store.set.assert_called_once_with("_iterable_completed", {"completed": False, "since": 2}) + + def test_attempt_after_a_failed_rerun_trusts_the_checkpoints_written_since(self): + context, store = self._context(try_number=3, marker={"completed": False, "since": 2}) + + with Checkpoints(context) as checkpoints: + assert checkpoints.trust_checkpoints is True + assert checkpoints.since == 2 + store.set.assert_not_called() + + store.set.assert_called_once_with("_iterable_completed", {"completed": True, "try_number": 3}) diff --git a/task-sdk/tests/task_sdk/definitions/test_iterate.py b/task-sdk/tests/task_sdk/definitions/test_iterate.py new file mode 100644 index 0000000000000..d808cc214dff0 --- /dev/null +++ b/task-sdk/tests/task_sdk/definitions/test_iterate.py @@ -0,0 +1,225 @@ +# +# 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 the ``.iterate()`` / ``.iterate_kwargs()`` entry points on partials and decorated tasks.""" + +from __future__ import annotations + +import contextlib +from collections.abc import Callable +from typing import Any +from unittest import mock + +import pytest + +from airflow.sdk import DAG, TaskInstanceState +from airflow.sdk.bases.xcom import BaseXCom +from airflow.sdk.execution_time.comms import ( + GetTICount, + GetXCom, + GetXComSequenceSlice, + SetXCom, + TICount, + XComResult, + XComSequenceSliceResult, +) + +RunTI = Callable[[DAG, str, int], TaskInstanceState] + + +class TestIterate: + def test_iterate_task_with_dict_return_annotation_pushes_whole_results( + self, run_ti: RunTI, mock_supervisor_comms + ): + """A Mapping return annotation makes @task infer multiple_outputs=True. The runner must not + apply that to an iterated task: its return value is the XComIterable aggregate, which is not a + dict, and every sub-task result is pushed whole rather than fanned out by key.""" + items = [{"dag_id": "a", "n": 1}, {"dag_id": "b", "n": 2}, {"dag_id": "c", "n": 3}] + + with DAG(dag_id="iterate_dict_return") as dag: + + @dag.task + def list_items(): + return items + + @dag.task + def enrich(item: dict) -> dict: + return {"dag_id": item["dag_id"], "n": item["n"] * 2} + + enrich.iterate(item=list_items()) + + assert enrich.multiple_outputs is True + + def mock_comms(msg): + if isinstance(msg, GetXCom): + if msg.task_id == "list_items": + return XComResult(key=BaseXCom.XCOM_RETURN_KEY, value=items) + elif isinstance(msg, GetXComSequenceSlice): + if msg.task_id == "list_items": + return XComSequenceSliceResult(root=items) + elif isinstance(msg, GetTICount): + return TICount(count=1) + return mock.DEFAULT + + mock_supervisor_comms.send.side_effect = mock_comms + + # Sub-task results are pushed through the async send, the aggregate through the sync one. + mock_supervisor_comms.asend.reset_mock() + assert run_ti(dag, "enrich", -1) == TaskInstanceState.SUCCESS + pushed: dict[str, Any] = { + msg.key: msg.value + for call in [*mock_supervisor_comms.send.mock_calls, *mock_supervisor_comms.asend.mock_calls] + if isinstance(msg := (call.kwargs.get("msg") or call.args[0]), SetXCom) + and msg.task_id == "enrich" + } + + assert pushed[BaseXCom.XCOM_RETURN_KEY]["__classname__"] == "airflow.sdk.bases.xcom.XComIterable" + assert pushed[BaseXCom.XCOM_RETURN_KEY]["__data__"]["map_index"] == -1 + assert not {"dag_id", "n"} & pushed.keys() + + sub_results = [ + value for key, value in pushed.items() if key.startswith(f"{BaseXCom.XCOM_RETURN_KEY}_") + ] + assert sorted(sub_results, key=lambda r: r["dag_id"]) == [ + {"dag_id": "a", "n": 2}, + {"dag_id": "b", "n": 4}, + {"dag_id": "c", "n": 6}, + ] + + def test_decorated_iterate_validates_as_iterate_not_expand(self): + """The decorated ``.iterate()`` names itself in its errors and, like ``.expand()``, refuses a + literal that is no collection, as the iteration refuses the same value from an upstream.""" + with DAG(dag_id="test_decorated_iterate_validation") as dag: + + @dag.task + def show(number): + return number + + with pytest.raises(TypeError, match=r"iterate\(\) got an unexpected keyword argument 'bogus'"): + show.iterate(bogus=1) + with pytest.raises(ValueError, match=r"cannot call iterate\(\) on task context variable 'ti'"): + show.iterate(ti=1) + with pytest.raises(ValueError, match=r"expand\(\) got an unexpected type 'int'"): + show.expand(number=5) + with pytest.raises(ValueError, match=r"iterate\(\) got an unexpected type 'int'"): + show.iterate(number=5) + with pytest.raises(ValueError, match=r"iterate\(\) got an unexpected type 'str'"): + show.iterate(number="abc") + with pytest.raises(ValueError, match=r"iterate\(\) got an unexpected type 'NoneType'"): + show.iterate(number=None) + + def test_iterate_marks_partial_as_expanded(self, recwarn): + """Test that .iterate() (like .expand()) flags the OperatorPartial as consumed, so + OperatorPartial.__del__ does not spuriously warn "Task ... was never mapped!" once the + partial and its resulting operator are garbage collected.""" + from airflow.providers.standard.operators.empty import EmptyOperator + from airflow.sdk.definitions._internal.expandinput import DictOfListsExpandInput + + with DAG(dag_id="test_iterate_expand_called"): + partial = EmptyOperator.partial(task_id="test_task") + partial._iterate(DictOfListsExpandInput({"retry_delay": [1, 2]}), strict=False) + + assert partial._expand_called is True + + del partial + assert not any("was never mapped" in str(w.message) for w in recwarn.list) + + +class TestIterateInTaskGroup: + """The iterated task and the operator it runs for each item have the same, once-prefixed task id.""" + + @staticmethod + def _dag(prefix_group_id: bool = True, nested: bool = False): + from airflow.sdk import BaseOperator, TaskGroup, task + + class Op(BaseOperator): + def __init__(self, x=None, **kwargs): + super().__init__(**kwargs) + self.x = x + + with DAG("in_task_group") as dag: + with TaskGroup("outer", prefix_group_id=prefix_group_id): + with TaskGroup("inner") if nested else contextlib.nullcontext(): + + @task + def f(x): + return x + + @task + def g(x): + return x + + decorated = f.iterate(x=[1, 2]).operator + g.expand(x=[1, 2]) + classic = Op.partial(task_id="c").iterate(x=[1, 2]) + return dag, decorated, classic + + @pytest.mark.parametrize( + ("prefix_group_id", "nested", "prefix"), + [ + pytest.param(True, False, "outer.", id="group"), + pytest.param(True, True, "outer.inner.", id="nested-groups"), + pytest.param(False, False, "", id="no-prefix"), + ], + ) + def test_task_ids_are_prefixed_once(self, prefix_group_id, nested, prefix): + dag, decorated, classic = self._dag(prefix_group_id, nested) + + assert sorted(dag.task_dict) == sorted(f"{prefix}{name}" for name in ("c", "f", "g")) + assert decorated.task_id == decorated._operator.task_id == f"{prefix}f" + assert classic.task_id == classic._operator.task_id == f"{prefix}c" + + +class TestIterateRejectsOperatorsThatSkipDownstream: + """ + An iteration has no downstream tasks of its own, so a skip-capable operator would skip nothing + and let every downstream task run: .iterate() refuses it when the Dag is defined. + """ + + @staticmethod + def _classic(name): + from airflow.providers.standard.operators.python import BranchPythonOperator, ShortCircuitOperator + + operator_class = {"short_circuit": ShortCircuitOperator, "branch": BranchPythonOperator}[name] + return operator_class.partial(task_id=name).iterate(python_callable=[lambda: True]) + + @staticmethod + def _decorated(name): + from airflow.sdk import task + + def decide(x): + return x + + return getattr(task, name)(decide).iterate(x=[1]) + + @pytest.mark.parametrize("name", ["short_circuit", "branch"]) + @pytest.mark.parametrize("build", ["_classic", "_decorated"]) + def test_skip_capable_operator_is_rejected(self, build, name): + with DAG(f"rejects_{build}_{name}"): + with pytest.raises(TypeError, match="can skip downstream tasks and cannot be iterated"): + getattr(self, build)(name) + + def test_operators_that_cannot_skip_are_accepted(self): + from airflow.providers.standard.operators.python import PythonOperator + from airflow.sdk import task + + def work(x): + return x + + with DAG("accepts"): + task(work).iterate(x=[1]) + PythonOperator.partial(task_id="classic").iterate(python_callable=[lambda: True]) diff --git a/task-sdk/tests/task_sdk/definitions/test_xcom_arg.py b/task-sdk/tests/task_sdk/definitions/test_xcom_arg.py index 2d1f8b2a4fa52..916f51d3817c9 100644 --- a/task-sdk/tests/task_sdk/definitions/test_xcom_arg.py +++ b/task-sdk/tests/task_sdk/definitions/test_xcom_arg.py @@ -17,17 +17,21 @@ # under the License. from __future__ import annotations +import asyncio +import threading from collections.abc import Callable from unittest import mock import pytest import structlog +from task_sdk.definitions.conftest import make_xcom_arg from airflow.sdk import TaskInstanceState from airflow.sdk.bases.xcom import BaseXCom +from airflow.sdk.definitions._internal.types import NOTSET from airflow.sdk.definitions.dag import DAG -from airflow.sdk.definitions.xcom_arg import PlainXComArg -from airflow.sdk.exceptions import AirflowSkipException +from airflow.sdk.definitions.xcom_arg import PlainXComArg, XComArg +from airflow.sdk.exceptions import AirflowSkipException, XComNotFound from airflow.sdk.execution_time.comms import GetXCom, XComResult, XComSequenceSliceResult from airflow.sdk.execution_time.lazy_sequence import LazyXComSequence from airflow.sdk.serde import deserialize, serialize @@ -416,3 +420,193 @@ def test_resolve_uses_xcom_pull_for_specific_index(self): assert resolved == "value-0" ti.xcom_pull.assert_called_once() assert ti.xcom_pull.call_args.kwargs["map_indexes"] == 0 + + @pytest.mark.asyncio + async def test_aresolve_stays_lazy_without_pulling(self): + arg = self._make_arg() + ti = self._make_ti(computed=None) + ti.axcom_pull = mock.AsyncMock() + + resolved = await arg.aresolve({"ti": ti}) + + assert isinstance(resolved, LazyXComSequence) + ti.axcom_pull.assert_not_awaited() + ti.xcom_pull.assert_not_called() + + @pytest.mark.asyncio + async def test_aresolve_uses_axcom_pull_for_specific_index(self): + arg = self._make_arg() + ti = self._make_ti(computed=0) + ti.axcom_pull = mock.AsyncMock(return_value="value-0") + + resolved = await arg.aresolve({"ti": ti}) + + assert resolved == "value-0" + ti.axcom_pull.assert_awaited_once_with( + task_ids="do_something", key="test", default=NOTSET, map_indexes=0 + ) + ti.xcom_pull.assert_not_called() + + @pytest.mark.asyncio + async def test_aresolve_computes_map_indexes_off_the_loop_thread(self): + """ + Counting upstream task instances is a blocking supervisor call; on the loop thread it would + deadlock with the ``asend`` calls of the iterated tasks in flight, so it must run in a worker. + """ + arg = self._make_arg() + ti = self._make_ti(computed=0) + ti.axcom_pull = mock.AsyncMock(return_value="value-0") + threads: list[threading.Thread] = [] + + def compute(**kwargs): + threads.append(threading.current_thread()) + return 0 + + ti.get_relevant_upstream_map_indexes.side_effect = compute + + await arg.aresolve({"ti": ti}) + + assert threads + assert threads[0] is not threading.current_thread() + + +class TestPlainXComArgAresolve: + """``aresolve`` on an unmapped upstream pulls through ``ti.axcom_pull`` and never ``ti.xcom_pull``.""" + + @staticmethod + def _make_arg(key: str = BaseXCom.XCOM_RETURN_KEY, multiple_outputs: bool = False): + operator = mock.MagicMock() + operator.is_mapped = False + operator.task_id = "do_something" + operator.dag_id = "test_dag" + operator.multiple_outputs = multiple_outputs + operator.get_closest_mapped_task_group.return_value = None + return PlainXComArg(operator=operator, key=key) + + @staticmethod + def _make_ti(pulled): + ti = mock.MagicMock() + ti.dag_id = "test_dag" + ti.xcom_pull.return_value = pulled + ti.axcom_pull = mock.AsyncMock(return_value=pulled) + return ti + + @pytest.mark.asyncio + async def test_aresolve_pulls_unmapped_instance(self): + arg = self._make_arg() + ti = self._make_ti(pulled=[1, 2, 3]) + + assert await arg.aresolve({"ti": ti}) == [1, 2, 3] + + ti.axcom_pull.assert_awaited_once_with( + task_ids="do_something", key=BaseXCom.XCOM_RETURN_KEY, default=NOTSET, map_indexes=None + ) + ti.xcom_pull.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("key", "multiple_outputs"), + [ + pytest.param(BaseXCom.XCOM_RETURN_KEY, False, id="return-value"), + pytest.param("other", True, id="multiple-outputs"), + ], + ) + async def test_aresolve_missing_xcom_gives_none_like_resolve(self, key, multiple_outputs): + arg = self._make_arg(key=key, multiple_outputs=multiple_outputs) + ti = self._make_ti(pulled=NOTSET) + + assert await arg.aresolve({"ti": ti}) is None + assert arg.resolve({"ti": ti}) is None + + @pytest.mark.asyncio + async def test_aresolve_missing_custom_key_raises_like_resolve(self): + arg = self._make_arg(key="other") + ti = self._make_ti(pulled=NOTSET) + + with pytest.raises(XComNotFound): + await arg.aresolve({"ti": ti}) + with pytest.raises(XComNotFound): + arg.resolve({"ti": ti}) + + +def _fail_sync_resolve(*args, **kwargs): + pytest.fail("synchronous resolve() must not be used on the async path") + + +class AsyncOnlyValues: + """A resolved value that can only be iterated asynchronously, like an XComIterable on the loop.""" + + def __init__(self, values): + self.values = values + + async def __aiter__(self): + for value in self.values: + yield value + + +class CustomXComArg(XComArg): + """A third-party subclass that only implements the two required methods.""" + + def __init__(self, values): + self.values = values + self.threads: list[threading.Thread] = [] + + def iter_references(self): + yield from () + + def resolve(self, context): + self.threads.append(threading.current_thread()) + return self.values + + +class TestXComArgSubclassDefaults: + """A subclass implementing only ``iter_references`` and ``resolve`` works on the iterated path.""" + + @pytest.mark.asyncio + async def test_aresolve_defaults_to_resolve_off_the_loop_thread(self): + arg = CustomXComArg([1, 2, 3]) + + assert await arg.aresolve({}) == [1, 2, 3] + assert arg.threads + assert all(thread is not threading.current_thread() for thread in arg.threads) + + +class TestXComArg: + @pytest.mark.asyncio + async def test_map_xcomarg_aresolve(self): + base = make_xcom_arg([1, 2, 3]) + base.resolve = _fail_sync_resolve + mapped = base.map(lambda x: x * 10) + assert list(await mapped.aresolve({})) == [10, 20, 30] + + @pytest.mark.asyncio + async def test_zip_xcomarg_aresolve(self): + a = make_xcom_arg([1, 2]) + b = make_xcom_arg([10, 20]) + a.resolve = b.resolve = _fail_sync_resolve + zipped = a.zip(b) + assert list(await zipped.aresolve({})) == [(1, 10), (2, 20)] + + @pytest.mark.asyncio + async def test_concat_xcomarg_aresolve(self): + a = make_xcom_arg([1, 2]) + b = make_xcom_arg([10, 20]) + a.resolve = b.resolve = _fail_sync_resolve + concatenated = a.concat(b) + assert list(await concatenated.aresolve({})) == [1, 2, 10, 20] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "build", [lambda a: a.map(lambda x: x), lambda a: a.zip(a), lambda a: a.concat(a)] + ) + async def test_composite_xcomarg_aresolve_rejects_non_sequence_like_resolve(self, build): + base = make_xcom_arg(42) + composite = build(base) + with pytest.raises(ValueError, match="expects sequence or dict"): + await composite.aresolve({}) + with pytest.raises(ValueError, match="expects sequence or dict"): + composite.resolve({}) + + def test_base_aresolve_falls_back_to_resolve(self): + """The base default is resolve() in a worker thread; PlainXComArg's axcom_pull override is an optimisation.""" + assert asyncio.run(super(PlainXComArg, make_xcom_arg([1, 2])).aresolve({})) == [1, 2] diff --git a/task-sdk/tests/task_sdk/execution_time/conftest.py b/task-sdk/tests/task_sdk/execution_time/conftest.py index 2d60b635f63f3..ece0433abe378 100644 --- a/task-sdk/tests/task_sdk/execution_time/conftest.py +++ b/task-sdk/tests/task_sdk/execution_time/conftest.py @@ -18,9 +18,16 @@ from __future__ import annotations import sys +from datetime import timedelta from socket import socketpair +from unittest import mock import pytest +from uuid6 import uuid7 + +from airflow.sdk import BaseAsyncOperator, BaseOperator, timezone +from airflow.sdk.api.datamodels._generated import TaskInstanceState +from airflow.sdk.execution_time.task_runner import IndexedTaskInstance @pytest.fixture @@ -43,3 +50,69 @@ def disable_capturing(): sys.stderr = sys.__stderr__ yield sys.stdin, sys.stdout, sys.stderr = old_in, old_out, old_err + + +@pytest.fixture +def make_indexed_ti(): + """Factory for creating IndexedTaskInstance objects for testing.""" + + def _make_indexed_ti( + *, + task_id: str = "my_task", + dag_id: str = "my_dag", + run_id: str = "run_1", + map_index: int = -1, + index: int | None = 0, + try_number: int = 0, + max_tries: int = 3, + is_async: bool = False, + retry_delay: timedelta = None, + retry_exponential_backoff: bool = False, + max_retry_delay: timedelta | None = None, + end_date=None, + start_date=None, + logical_date=None, + do_xcom_push: bool = True, + ) -> IndexedTaskInstance: + """Create a IndexedTaskInstance via model_construct to bypass full Pydantic validation.""" + if retry_delay is None: + retry_delay = timedelta(seconds=300) + + # Set defaults for dates if not provided + if end_date is None: + end_date = timezone.datetime(2024, 12, 3, 10, 0, 0) + if start_date is None: + start_date = timezone.datetime(2024, 12, 3, 9, 55, 0) + if logical_date is None: + logical_date = timezone.datetime(2024, 12, 3, 0, 0, 0) + + operator_cls = BaseAsyncOperator if is_async else BaseOperator + operator = mock.create_autospec(operator_cls, instance=True) + operator.task_id = task_id + operator.dag_id = dag_id + operator.is_async = is_async + operator.retries = max_tries + operator.retry_delay = retry_delay + operator.retry_exponential_backoff = retry_exponential_backoff + operator.max_retry_delay = max_retry_delay + operator.do_xcom_push = do_xcom_push + + return IndexedTaskInstance.model_construct( + id=uuid7(), + task_id=task_id, + dag_id=dag_id, + run_id=run_id, + map_index=map_index, + index=index, + try_number=try_number, + max_tries=max_tries, + state=TaskInstanceState.SCHEDULED, + is_mapped=True, + task=operator, + dag_version_id=uuid7(), + end_date=end_date, + start_date=start_date, + logical_date=logical_date, + ) + + return _make_indexed_ti diff --git a/task-sdk/tests/task_sdk/execution_time/test_comms.py b/task-sdk/tests/task_sdk/execution_time/test_comms.py index f328380106ae5..514f182d79267 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_comms.py +++ b/task-sdk/tests/task_sdk/execution_time/test_comms.py @@ -246,6 +246,45 @@ def _hold_lock(): lock_release.set() holder.join(timeout=2) + @pytest.mark.asyncio + async def test_asend_cancelled_while_waiting_for_the_thread_lock_leaves_it_free(self, socket_pair): + """ + asend() waits for _thread_lock in a worker thread. Cancelling that wait (the executor shutting + down after a failure) cancels the future, not the thread: once the holder lets go, the thread + acquires, and the lock must not stay taken with nobody to release it. + """ + r, _ = socket_pair + decoder = CommsDecoder(socket=r, log=structlog.get_logger()) + + lock_held = threading.Event() + lock_release = threading.Event() + + def _hold_lock(): + decoder._thread_lock.acquire() + lock_held.set() + lock_release.wait() + decoder._thread_lock.release() + + holder = threading.Thread(target=_hold_lock, daemon=True) + holder.start() + assert lock_held.wait(timeout=2), "Background thread never acquired _thread_lock" + + parked = asyncio.ensure_future(decoder.asend(GetVariable(key="parked"))) + await asyncio.sleep(0.05) + parked.cancel() + with pytest.raises(asyncio.CancelledError): + await parked + + lock_release.set() + holder.join(timeout=2) + for _ in range(200): + if decoder._thread_lock.acquire(blocking=False): + decoder._thread_lock.release() + break + await asyncio.sleep(0.01) + else: + pytest.fail("_thread_lock stayed taken after the cancelled asend()") + @pytest.mark.asyncio async def test_send_from_event_loop_succeeds_when_lock_free(self, socket_pair): """ diff --git a/task-sdk/tests/task_sdk/execution_time/test_context.py b/task-sdk/tests/task_sdk/execution_time/test_context.py index 2d96f03e94828..fed59bac310f1 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_context.py +++ b/task-sdk/tests/task_sdk/execution_time/test_context.py @@ -17,7 +17,11 @@ from __future__ import annotations +import asyncio +import contextvars import json +import threading +from concurrent.futures import ThreadPoolExecutor from datetime import UTC, datetime, timedelta from typing import TYPE_CHECKING from unittest import mock @@ -96,6 +100,7 @@ AssetStateStoreAccessor, AssetStateStoreAccessors, ConnectionAccessor, + IndexedTaskStateStoreAccessor, InletEventsAccessors, MacrosAccessor, OutletEventAccessor, @@ -115,6 +120,7 @@ _wrap_external_ref, context_to_airflow_vars, set_current_context, + set_indexed_context, ) from airflow.sdk.execution_time.secrets import ExecutionAPISecretsBackend from airflow.sdk.state import BaseStoreBackend @@ -569,6 +575,100 @@ def test_nested_context(self): # End of with statement ctx_list[i].__exit__(None, None, None) + @pytest.mark.parametrize("start", ["thread_pool", "thread"]) + def test_thread_started_by_the_task_sees_its_context(self, start): + """A thread the task starts has empty ContextVars, and must still find the task's context.""" + task_context = {"Hello": "World"} + + def read(): + return get_current_context() + + with set_current_context(task_context): + if start == "thread_pool": + with ThreadPoolExecutor(max_workers=1) as pool: + seen = pool.submit(read).result() + else: + results = [] + thread = threading.Thread(target=lambda: results.append(read())) + thread.start() + thread.join() + (seen,) = results + + assert seen is task_context + + def test_indexed_context_covers_the_task_context_within_its_block(self): + task_context = {"ContextId": "task"} + indexed_context = {"ContextId": "iteration"} + + with set_current_context(task_context): + with set_indexed_context(indexed_context): + assert get_current_context() is indexed_context + assert get_current_context() is task_context + + def test_indexed_context_is_not_seen_by_other_threads(self): + """A thread started inside an iteration sees the task's context, not the iteration's.""" + task_context = {"ContextId": "task"} + + with set_current_context(task_context): + with set_indexed_context({"ContextId": "iteration"}): + with ThreadPoolExecutor(max_workers=1) as pool: + seen = pool.submit(get_current_context).result() + + assert seen is task_context + + @pytest.mark.asyncio + async def test_helpers_that_carry_the_iterations_context_over(self): + """ + What the docs point iterations at for helper threads: ``asyncio.to_thread`` and + ``copy_context().run`` see the iteration's context, ``run_in_executor`` the task's. + """ + task_context = {"ContextId": "task"} + indexed_context = {"ContextId": "iteration"} + + with set_current_context(task_context): + with set_indexed_context(indexed_context): + in_to_thread = await asyncio.to_thread(get_current_context) + with ThreadPoolExecutor(max_workers=1) as pool: + copied = contextvars.copy_context() + in_copied_context = pool.submit(copied.run, get_current_context).result() + in_executor = await asyncio.get_running_loop().run_in_executor(pool, get_current_context) + + assert in_to_thread is indexed_context + assert in_copied_context is indexed_context + assert in_executor is task_context + + @pytest.mark.asyncio + async def test_concurrent_iterations_each_see_their_own_context(self): + """ + Iterations interleaving on one event loop never see each other's context. + + Each one reads its context while the other is inside its own block, and neither leaves + before both have read, so a stack shared by the thread (as a thread-local would be) hands + one of them the other's context. + """ + entered: list[int] = [] + seen: dict[int, object] = {} + both_entered = asyncio.Event() + both_read = asyncio.Event() + + async def iteration(index): + with set_indexed_context({"ContextId": index}): + entered.append(index) + if len(entered) == 2: + both_entered.set() + await both_entered.wait() + # Let the other iteration resume inside its block before reading. + await asyncio.sleep(0) + seen[index] = get_current_context()["ContextId"] + if len(seen) == 2: + both_read.set() + await both_read.wait() + + with set_current_context({"ContextId": "task"}): + await asyncio.gather(iteration(0), iteration(1)) + assert seen == {0: 0, 1: 1} + assert get_current_context()["ContextId"] == "task" + class TestOutletEventAccessor: @pytest.mark.parametrize( @@ -957,6 +1057,32 @@ def test_for_asset_alias(self, mocked__getitem__): outlet_event_accessors.for_asset_alias(name="name") assert mocked__getitem__.call_args[0][0] == TEST_ASSET_ALIAS + def test_concurrent_access_same_asset_preserves_accessor(self): + """Concurrent __getitem__ for the same asset must not overwrite an existing accessor.""" + import threading + + accessors = OutletEventAccessors() + asset = Asset("concurrent-test") + results: list[OutletEventAccessor] = [] + + barrier = threading.Barrier(2) + + def access(): + barrier.wait() + results.append(accessors[asset]) + + threads = [threading.Thread(target=access) for _ in range(2)] + for t in threads: + t.start() + for t in threads: + t.join() + + assert len(results) == 2 + # Both threads must have received the identical accessor object so + # that neither thread's accumulated events can be silently discarded. + assert results[0] is results[1] + assert len(accessors) == 1 + class TestInletEventAccessor: @pytest.fixture @@ -2960,3 +3086,87 @@ def test_nothing_is_hidden_when_multi_team_is_off(self, integrated_macros): assert accessor.team_a_macros.team_a_macro() == "team-a" assert accessor.global_macros.shared_macro() == "shared" + + +class TestIndexedTaskStateStoreAccessor: + """The parent's store seen from one iteration: every key carries the index, clearing is refused.""" + + @pytest.fixture + def store(self): + return mock.create_autospec(TaskStateStoreAccessor, instance=True) + + def test_get_and_set_suffix_the_key(self, store): + indexed = IndexedTaskStateStoreAccessor(store, index=2) + store.get.return_value = 42 + + indexed.set("last_offset", 42, retention=timedelta(hours=1)) + assert indexed.get("last_offset", default=0) == 42 + + store.set.assert_called_once_with("last_offset_2", 42, retention=timedelta(hours=1)) + store.get.assert_called_once_with("last_offset_2", 0) + + def test_delete_suffixes_the_key(self, store): + IndexedTaskStateStoreAccessor(store, index=0).delete("last_offset") + store.delete.assert_called_once_with("last_offset_0") + + @pytest.mark.asyncio + async def test_async_reads_and_writes_suffix_the_key(self, store): + indexed = IndexedTaskStateStoreAccessor(store, index=7) + store.aget.return_value = "x" + + await indexed.aset("cursor", "x") + assert await indexed.aget("cursor") == "x" + await indexed.adelete("cursor") + + store.aset.assert_awaited_once_with("cursor_7", "x", retention=None) + store.aget.assert_awaited_once_with("cursor_7", None) + store.adelete.assert_awaited_once_with("cursor_7") + + @pytest.mark.asyncio + async def test_clear_is_refused(self, store): + indexed = IndexedTaskStateStoreAccessor(store, index=1) + + with pytest.raises(RuntimeError, match="not available inside an iterated task"): + indexed.clear() + with pytest.raises(RuntimeError, match="not available inside an iterated task"): + await indexed.aclear() + store.clear.assert_not_called() + store.aclear.assert_not_called() + + def test_the_runners_backend_clear_is_refused_as_well(self, store): + """The view has no scope of its own to clear; the inherited path is refused, not broken.""" + with pytest.raises(RuntimeError, match="not available inside an iterated task"): + IndexedTaskStateStoreAccessor(store, index=1)._clear_backend_only() + + def test_the_view_never_equals_the_parents_accessor(self): + """Python tries the subclass's ``__eq__`` first, so the parent's never reads the view's ``_ti_id``.""" + parent = TaskStateStoreAccessor(UUID(int=1), TaskScope(dag_id="d", run_id="r", task_id="t")) + indexed = IndexedTaskStateStoreAccessor(parent, index=1) + + assert parent != indexed + assert indexed != parent + assert indexed == IndexedTaskStateStoreAccessor(parent, index=1) + + @pytest.mark.parametrize("key", ["_iterable", "_iterable_completed", "_iterable_3"]) + @pytest.mark.asyncio + async def test_keys_of_the_operators_checkpoints_are_refused(self, store, key): + """``_iterable`` written from iteration 0 would land on ``_iterable_0``, that index's checkpoint.""" + indexed = IndexedTaskStateStoreAccessor(store, index=0) + + for call in ( + lambda: indexed.get(key), + lambda: indexed.set(key, 1), + lambda: indexed.delete(key), + ): + with pytest.raises(ValueError, match="reserved for the checkpoints"): + call() + for acall in (lambda: indexed.aget(key), lambda: indexed.aset(key, 1), lambda: indexed.adelete(key)): + with pytest.raises(ValueError, match="reserved for the checkpoints"): + await acall() + assert store.mock_calls == [] + + def test_identity_is_the_store_and_the_index(self, store): + assert IndexedTaskStateStoreAccessor(store, index=1) == IndexedTaskStateStoreAccessor(store, index=1) + assert IndexedTaskStateStoreAccessor(store, index=1) != IndexedTaskStateStoreAccessor(store, index=2) + assert IndexedTaskStateStoreAccessor(store, index=1) != store + assert "index=1" in repr(IndexedTaskStateStoreAccessor(store, index=1)) diff --git a/task-sdk/tests/task_sdk/execution_time/test_executor.py b/task-sdk/tests/task_sdk/execution_time/test_executor.py new file mode 100644 index 0000000000000..fce5fad0f7598 --- /dev/null +++ b/task-sdk/tests/task_sdk/execution_time/test_executor.py @@ -0,0 +1,339 @@ +# 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 + +import asyncio +import threading +import time +from concurrent.futures import Executor +from unittest import mock + +import pytest + +from airflow.sdk.bases.operator import event_loop +from airflow.sdk.execution_time.executor import AsyncAwareExecutor + + +async def aiter(items): + """Expose a list as the async iterable ``AsyncAwareExecutor.imap_unordered`` consumes.""" + for item in items: + yield item + + +class TestAsyncAwareExecutor: + def test_submit_sync_function_returns_future(self): + """Sync callables are dispatched to the thread pool and return an asyncio.Future.""" + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + future = executor.submit(lambda: 42) + assert isinstance(future, asyncio.Future) + assert loop.run_until_complete(future) == 42 + + def test_submit_async_coroutine_function_returns_task(self): + """Async callables are scheduled on the event loop and return an asyncio.Task.""" + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + + async def async_fn(): + return "async_result" + + task = executor.submit(async_fn) + assert isinstance(task, asyncio.Task) + result = loop.run_until_complete(task) + assert result == "async_result" + + def test_submit_coroutine_object_returns_task(self): + """Passing a coroutine object (not a function) directly is also scheduled on the event loop.""" + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + + async def async_fn(): + return "coro_result" + + coro = async_fn() + task = executor.submit(coro) + assert isinstance(task, asyncio.Task) + result = loop.run_until_complete(task) + assert result == "coro_result" + + def test_submit_sync_function_propagates_exception(self): + """Exceptions raised inside sync callables are propagated when the future is resolved.""" + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + future = executor.submit(lambda: (_ for _ in ()).throw(ValueError("boom"))) + assert isinstance(future, asyncio.Future) + with pytest.raises(ValueError, match="boom"): + loop.run_until_complete(future) + + def test_submit_async_function_propagates_exception(self): + """Exceptions raised inside async callables are propagated when the task is awaited.""" + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + + async def failing(): + raise RuntimeError("async boom") + + task = executor.submit(failing) + with pytest.raises(RuntimeError, match="async boom"): + loop.run_until_complete(task) + + def test_semaphore_limits_concurrent_async_tasks(self): + """The semaphore prevents more than max_workers coroutines from running simultaneously.""" + concurrency_high_watermark = 0 + running = 0 + + async def count_concurrent(): + nonlocal concurrency_high_watermark, running + running += 1 + concurrency_high_watermark = max(concurrency_high_watermark, running) + await asyncio.sleep(0) + running -= 1 + + max_workers = 2 + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=max_workers) as executor: + tasks = [executor.submit(count_concurrent) for _ in range(6)] + loop.run_until_complete(asyncio.gather(*tasks)) + + assert concurrency_high_watermark <= max_workers + + def test_exit_shuts_down_thread_pool(self): + """__exit__ calls shutdown on the thread pool without blocking indefinitely on it.""" + with event_loop() as loop: + executor = AsyncAwareExecutor(loop=loop, max_workers=2) + with mock.patch.object( + executor._thread_pool, "shutdown", wraps=executor._thread_pool.shutdown + ) as shutdown_mock: + with executor: + pass + # wait=False: the executor itself bounds how long it waits on worker + # threads afterwards instead of delegating an unbounded wait to the pool. + shutdown_mock.assert_called_once_with(wait=False, cancel_futures=False) + + def test_context_manager_returns_self(self): + """__enter__ returns the executor instance itself.""" + with event_loop() as loop: + executor = AsyncAwareExecutor(loop=loop, max_workers=2) + with executor as ctx: + assert ctx is executor + + def test_does_not_override_executor_map(self): + """Completion-order streaming is imap_unordered; Executor.map keeps its submission-order contract.""" + assert AsyncAwareExecutor.map is Executor.map + + def test_imap_unordered_streams_completed_sync_results(self): + """map() yields completed results as work finishes instead of waiting for all items.""" + + def sleepy_value(delay: float) -> float: + time.sleep(delay) + return delay + + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + started = time.monotonic() + result_iter = executor.imap_unordered(sleepy_value, aiter([0.25, 0.01])) + first = next(result_iter) + + assert first == 0.01 + # The faster work (0.01s) should complete well before the slower work (0.25s). + # Allow overhead for thread pool scheduling (typically ~0.15-0.2s on busy systems). + assert time.monotonic() - started < 0.35 + + def test_imap_unordered_stops_pulling_once_stop_says_so(self): + """Once ``stop`` answers True nothing more is pulled or submitted; what was submitted drains.""" + pulled: list[int] = [] + results: list[int] = [] + + async def source(): + for item in range(6): + pulled.append(item) + yield item + + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + for result in executor.imap_unordered(lambda x: x, source(), stop=lambda: len(results) >= 1): + results.append(result) + + assert len(pulled) <= 3 + assert sorted(results) == sorted(pulled) + + def test_shutdown_cancel_futures_cancels_async_tasks(self): + """shutdown(cancel_futures=True) cancels submitted async tasks.""" + + async def long_running(): + await asyncio.sleep(60) + + with event_loop() as loop: + executor = AsyncAwareExecutor(loop=loop, max_workers=2) + task = executor.submit(long_running) + + executor.shutdown(wait=False, cancel_futures=True) + loop.run_until_complete(asyncio.sleep(0)) + + assert task.cancelled() + + def test_submit_after_shutdown_raises_runtime_error(self): + with event_loop() as loop: + executor = AsyncAwareExecutor(loop=loop, max_workers=2) + executor.shutdown(wait=False) + + with pytest.raises(RuntimeError, match="cannot schedule new futures after shutdown"): + executor.submit(lambda: 1) + + def test_imap_unordered_zips_async_iterables_and_stops_at_the_shortest(self): + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + results = sorted( + executor.imap_unordered(lambda a, b: (a, b), aiter([1, 2, 3]), aiter([10, 20])) + ) + + assert results == [(1, 10), (2, 20)] + + def test_imap_unordered_pulls_items_on_the_running_loop_and_lazily(self): + """ + The next item is pulled from a coroutine while the loop runs, and only when a slot frees up. + + Both are the guard against the deadlock of pulling from the main thread between two + ``run_until_complete`` calls: a pull that needs the supervisor channel then blocks on the lock + held by a call parked mid-``asend``, which can only complete once the loop runs again. + """ + pulled: list[int] = [] + loop_running_at_pull: list[bool] = [] + + async def source(): + for item in range(6): + loop_running_at_pull.append(asyncio.get_running_loop().is_running()) + pulled.append(item) + yield item + + async def work(item: int) -> int: + await asyncio.sleep(0.01) + return item + + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + result_iter = executor.imap_unordered(work, source()) + first = next(result_iter) + # Only max_workers items were pulled to start; the rest wait for free slots. + assert first in (0, 1) + assert pulled == [0, 1] or pulled == [0, 1, 2] + rest = list(result_iter) + + assert sorted([first, *rest]) == list(range(6)) + assert pulled == list(range(6)) + assert loop_running_at_pull == [True] * 6 + + def test_imap_unordered_does_not_deadlock_when_pulling_needs_a_lock_held_by_a_parked_call(self): + """ + Regression test for the IterableOperator freeze. + + A call holds a thread lock across an ``await`` (the supervisor channel's lock held by a + parked ``asend``), and pulling the next item takes that same lock from a worker thread (a + synchronous SDK read behind an iterated input). If the pull happened while the loop was + paused, the parked call could never release the lock and the process would freeze with + every thread idle. Pulling on the running loop lets the parked call finish first. + """ + lock = threading.Lock() + + async def source(): + for item in range(4): + await asyncio.to_thread(lock.acquire) + lock.release() + yield item + + async def work(item: int) -> int: + loop = asyncio.get_running_loop() + await loop.run_in_executor(None, lock.acquire) + try: + await asyncio.sleep(0.02) # parked while holding the lock + finally: + lock.release() + return item + + results: list[int] = [] + + def run() -> None: + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + results.extend(executor.imap_unordered(work, source())) + + thread = threading.Thread(target=run, daemon=True) + thread.start() + thread.join(timeout=10) + + assert not thread.is_alive(), "map() deadlocked while pulling the next item" + assert sorted(results) == [0, 1, 2, 3] + + def test_imap_unordered_timeout_raises_timeout_error(self): + def slow_fn(delay: float) -> float: + time.sleep(delay) + return delay + + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=1) as executor: + with pytest.raises(TimeoutError): + list(executor.imap_unordered(slow_fn, aiter([0.2]), timeout=0.01)) + + def test_imap_unordered_streams_completed_async_results(self): + async def async_sleepy_value(delay: float) -> float: + await asyncio.sleep(delay) + return delay + + with event_loop() as loop: + with AsyncAwareExecutor(loop=loop, max_workers=2) as executor: + started = time.monotonic() + result_iter = executor.imap_unordered(async_sleepy_value, aiter([0.2, 0.01])) + first = next(result_iter) + + assert first == 0.01 + # The faster work (0.01s) should complete well before the slower work (0.2s). + # Allow overhead for event loop scheduling (typically ~0.1-0.15s on busy systems). + assert time.monotonic() - started < 0.3 + + def test_shutdown_wait_true_waits_for_async_tasks(self): + async def short_running() -> str: + await asyncio.sleep(0.01) + return "done" + + with event_loop() as loop: + executor = AsyncAwareExecutor(loop=loop, max_workers=2) + task = executor.submit(short_running) + + executor.shutdown(wait=True) + + assert task.done() + assert not task.cancelled() + assert task.result() == "done" + + def test_shutdown_does_not_block_forever_on_stuck_worker_thread(self): + """shutdown(wait=True) must be bounded by shutdown_timeout, not hang on a stuck thread.""" + release = threading.Event() + + def blocking_fn(): + release.wait(timeout=5) + return "done" + + with event_loop() as loop: + executor = AsyncAwareExecutor(loop=loop, max_workers=1, shutdown_timeout=0.05) + executor.submit(blocking_fn) + + started = time.monotonic() + executor.shutdown(wait=True) + elapsed = time.monotonic() - started + + release.set() + assert elapsed < 1.0 diff --git a/task-sdk/tests/task_sdk/execution_time/test_lazy_sequence.py b/task-sdk/tests/task_sdk/execution_time/test_lazy_sequence.py index 16149ae9e26c3..b1e6f135d4e13 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_lazy_sequence.py +++ b/task-sdk/tests/task_sdk/execution_time/test_lazy_sequence.py @@ -17,7 +17,7 @@ from __future__ import annotations -from unittest.mock import Mock, call +from unittest.mock import AsyncMock, Mock, call import pytest @@ -59,6 +59,47 @@ def lazy_sequence(mock_xcom_arg, mock_ti): return LazyXComSequence(mock_xcom_arg, mock_ti) +@pytest.mark.asyncio +async def test_aget(mock_supervisor_comms, lazy_sequence): + mock_supervisor_comms.asend = AsyncMock(return_value=XComSequenceIndexResult(root="f")) + + assert await lazy_sequence.aget(1) == "f" + + mock_supervisor_comms.asend.assert_awaited_once_with( + GetXComSequenceItem( + key=BaseXCom.XCOM_RETURN_KEY, dag_id="dag", task_id="task", run_id="run", offset=1 + ), + ) + mock_supervisor_comms.send.assert_not_called() + + +@pytest.mark.asyncio +async def test_aget_out_of_range_raises_index_error(mock_supervisor_comms, lazy_sequence): + mock_supervisor_comms.asend = AsyncMock( + return_value=ErrorResponse(error=ErrorType.XCOM_NOT_FOUND, detail={"oops": "sorry!"}) + ) + + with pytest.raises(IndexError): + await lazy_sequence.aget(3) + + +@pytest.mark.asyncio +async def test_aiter(mock_supervisor_comms, lazy_sequence): + """``async for`` fetches item by item through ``asend`` and never through the blocking ``send``.""" + mock_supervisor_comms.asend = AsyncMock( + side_effect=[ + XComSequenceIndexResult(root="f"), + XComSequenceIndexResult(root="g"), + ErrorResponse(error=ErrorType.XCOM_NOT_FOUND, detail={"oops": "sorry!"}), + ] + ) + + assert [item async for item in lazy_sequence] == ["f", "g"] + + assert [call.args[0].offset for call in mock_supervisor_comms.asend.await_args_list] == [0, 1, 2] + mock_supervisor_comms.send.assert_not_called() + + class CustomXCom(BaseXCom): @classmethod def deserialize_value(cls, xcom): diff --git a/task-sdk/tests/task_sdk/execution_time/test_task_runner.py b/task-sdk/tests/task_sdk/execution_time/test_task_runner.py index c5ec20f36c6af..335e429b7fdd5 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_task_runner.py +++ b/task-sdk/tests/task_sdk/execution_time/test_task_runner.py @@ -17,6 +17,7 @@ from __future__ import annotations +import asyncio import contextlib import contextvars import functools @@ -82,7 +83,16 @@ from airflow.sdk.coordinators._dag_importer import CoordinatorDagImporter from airflow.sdk.coordinators._subprocess import SubprocessCoordinator from airflow.sdk.definitions._internal.types import NOTSET, SET_DURING_EXECUTION, is_arg_set -from airflow.sdk.definitions.asset import Asset, AssetAlias, AssetUniqueKey, AssetUriRef, Dataset, Model +from airflow.sdk.definitions.asset import ( + Asset, + AssetAlias, + AssetAliasEvent, + AssetUniqueKey, + AssetUriRef, + Dataset, + Model, +) +from airflow.sdk.definitions.iterableoperator import IterationState from airflow.sdk.definitions.param import DagParam from airflow.sdk.definitions.retry_policy import ( ExceptionRetryPolicy, @@ -167,6 +177,7 @@ ) from airflow.sdk.execution_time.context import ( ConnectionAccessor, + IndexedTaskStateStoreAccessor, InletEventsAccessors, MacrosAccessor, OutletEventAccessors, @@ -176,6 +187,9 @@ _wrap_external_ref, ) from airflow.sdk.execution_time.task_runner import ( + IndexedTaskInstance, + IndexedTaskRunner, + IndexedTaskState, RuntimeTaskInstance, TaskRunnerMarker, _defer_task, @@ -201,6 +215,7 @@ from airflow.triggers.testing import SuccessTrigger from tests_common.test_utils.config import conf_vars +from tests_common.test_utils.mock_context import mock_context from tests_common.test_utils.mock_operators import AirflowLink if TYPE_CHECKING: @@ -1848,6 +1863,36 @@ def sleep_and_catch_other_exceptions(): _execute_task(context=ti.get_template_context(), ti=ti, log=mock.MagicMock()) +def test_execution_timeout_caps_iterated_task_with_sync_sub_tasks(create_runtime_ti, mock_supervisor_comms): + """The wrapped operator's execution_timeout is kept on the IterableOperator as the wall-clock cap on + the whole iteration. The runner enforces it on the main thread, so it fires even though the sync + sub-tasks run in worker threads where SIGALRM cannot reach them.""" + from airflow.sdk.definitions._internal.expandinput import ListOfDictsExpandInput + from airflow.sdk.definitions.iterableoperator import IterableOperator + + class SleepyOperator(BaseOperator): + def execute(self, context): + time.sleep(1) + + # Two items on a single worker: without the cap the iteration takes two sleeps. With it the + # timeout fires during the first item, and the second is never started. Python threads cannot + # be interrupted, so the error surfaces once the first item's sleep ends, after one sleep. + expand_input = ListOfDictsExpandInput([{}, {}]) + with DAG("dag_iterate_execution_timeout") as dag: + mapped = SleepyOperator.partial( + task_id="sleepy", dag=dag, task_concurrency=1, execution_timeout=timedelta(milliseconds=200) + )._expand(expand_input, strict=True, register_with_dag=False) + op = IterableOperator(operator=mapped, expand_input=expand_input, dag=dag) + assert op.execution_timeout == timedelta(milliseconds=200) + + ti = create_runtime_ti(task=op, dag_id="dag_iterate_execution_timeout") + + started = time.monotonic() + with pytest.raises(AirflowTaskTimeout): + _execute_task(context=ti.get_template_context(), ti=ti, log=mock.MagicMock()) + assert time.monotonic() - started < 1.8 + + def test_basic_templated_dag(mocked_parse, make_ti_context, mock_supervisor_comms, spy_agency): """Test running a Dag with templated task.""" from airflow.providers.standard.operators.bash import BashOperator @@ -2467,29 +2512,46 @@ def test_run_with_asset_inlets(create_runtime_ti, mock_supervisor_comms): inlet_events[Asset(name="no such asset in inlets")] -@mock.patch("airflow.sdk.execution_time.task_runner.context_to_airflow_vars") @mock.patch.dict(os.environ, {}, clear=True) -def test_execute_task_exports_env_vars( - mock_context_to_airflow_vars, create_runtime_ti, mock_supervisor_comms +def test_execute_task_exports_context_vars_to_environ(create_runtime_ti, mock_supervisor_comms): + """A regular task instance exports AIRFLOW_CTX_* to os.environ before the operator runs.""" + captured_vars = {} + + def test_function(): + captured_vars["dag_id"] = os.environ.get("AIRFLOW_CTX_DAG_ID") + captured_vars["task_id"] = os.environ.get("AIRFLOW_CTX_TASK_ID") + return "test function" + + task = PythonOperator(task_id="test_task", python_callable=test_function) + ti = create_runtime_ti(task=task, dag_id="dag_with_ctx_vars") + + run(ti, ti.get_template_context(), log=mock.MagicMock()) + + assert captured_vars == {"dag_id": "dag_with_ctx_vars", "task_id": "test_task"} + + +@mock.patch.dict(os.environ, {}, clear=True) +def test_execute_task_leaves_environ_alone_for_indexed_task_instance( + create_runtime_ti, mock_supervisor_comms ): - """Test that _execute_task exports airflow context to environment variables.""" + """Indexed sub-tasks run concurrently in one process, so _execute_task must not touch the shared + os.environ for them; the parent IterableOperator exported the same values before they started.""" + captured_vars = {} def test_function(): + captured_vars["dag_id"] = os.environ.get("AIRFLOW_CTX_DAG_ID") return "test function" - task = PythonOperator( - task_id="test_task", - python_callable=test_function, + task = PythonOperator(task_id="test_task", python_callable=test_function) + ti = create_runtime_ti(task=task, dag_id="dag_with_indexed_ctx_vars") + indexed_ti = IndexedTaskInstance.create_indexed_task( + context={"ti": ti, "task": task, "task_state_store": ti.task_state_store}, index=0, operator=task ) - ti = create_runtime_ti(task=task, dag_id="dag_with_env_vars") + _execute_task(context=indexed_ti.get_template_context(), ti=indexed_ti, log=mock.MagicMock()) - mock_env_vars = {"AIRFLOW_CTX_DAG_ID": "test_dag_env_vars", "AIRFLOW_CTX_TASK_ID": "test_env_task"} - mock_context_to_airflow_vars.return_value = mock_env_vars - run(ti, ti.get_template_context(), log=mock.MagicMock()) - - assert os.environ["AIRFLOW_CTX_DAG_ID"] == "test_dag_env_vars" - assert os.environ["AIRFLOW_CTX_TASK_ID"] == "test_env_task" + assert captured_vars == {"dag_id": None} + assert "AIRFLOW_CTX_DAG_ID" not in os.environ def test_execute_success_task_with_rendered_map_index(create_runtime_ti, mock_supervisor_comms): @@ -2545,6 +2607,785 @@ def test_function(ti): assert ti.rendered_map_index == "Label: test_task" +class TestIndexedTaskState: + @pytest.mark.parametrize( + "result", + [ + pytest.param(('[{"@odata.context": "..."}]', "6C5EF9E6CBBE"), id="tuple"), + pytest.param(datetime(2026, 9, 17, 21, 16, 7, tzinfo=UTC), id="datetime"), + pytest.param([{"a": 1}, {"b": [1, 2]}], id="plain_json"), + ], + ) + def test_result_survives_the_state_store_message_and_round_trips(self, result): + """The checkpoint is sent as a JsonValue, so a result that is not plain JSON has to be + serialized on the way in and restored on the way out.""" + serialized = IndexedTaskState(status=TaskInstanceState.SUCCESS, result=result).serialize() + + msg = SetTaskStateStore(ti_id=uuid7(), key="_iterable_0", value=serialized, expires_at=None) + + assert IndexedTaskState.deserialize(msg.value).result == result + + def test_extra_xcoms_survive_the_state_store_message_and_round_trip(self): + """Values an item pushed are serialized like its result, since the checkpoint travels as JSON.""" + xcoms = { + "when": datetime(2026, 9, 30, 11, 0, tzinfo=UTC), + "pair": ("a", 1), + "plain": {"rows": [1, 2]}, + } + serialized = IndexedTaskState(status=TaskInstanceState.SUCCESS, xcoms=xcoms).serialize() + + msg = SetTaskStateStore(ti_id=uuid7(), key="_iterable_0", value=serialized, expires_at=None) + + assert IndexedTaskState.deserialize(msg.value).xcoms == xcoms + + def test_a_checkpoint_without_extra_xcoms_leaves_them_out(self): + serialized = IndexedTaskState(status=TaskInstanceState.SUCCESS).serialize() + + assert "xcoms" not in serialized + assert IndexedTaskState.deserialize(serialized).xcoms is None + + @staticmethod + def _recorded_accessors() -> OutletEventAccessors: + accessors = OutletEventAccessors() + asset = accessors[Asset(name="a", uri="s3://bucket/a")] + asset.extra = {"rows": 3} + asset.add_partitions(["p2", "p1"]) + accessors[AssetAlias(name="alias")].asset_alias_events.append( + AssetAliasEvent( + source_alias_name="alias", + dest_asset_key=AssetUniqueKey(name="b", uri="s3://bucket/b"), + dest_asset_extra={"via": "alias"}, + extra={"n": 1}, + ) + ) + return accessors + + def test_record_outlet_events_keeps_a_json_safe_snapshot(self): + """The snapshot travels with the checkpoint, so it holds plain JSON, partition keys sorted.""" + state = IndexedTaskState(status=TaskInstanceState.SUCCESS) + + state.record_outlet_events(self._recorded_accessors()) + + assert state.outlet_events == [ + { + "kind": "asset", + "name": "a", + "uri": "s3://bucket/a", + "extra": {"rows": 3}, + "partition_keys": ["p1", "p2"], + }, + { + "kind": "asset_alias", + "source_alias_name": "alias", + "dest_asset_key": {"name": "b", "uri": "s3://bucket/b"}, + "dest_asset_extra": {"via": "alias"}, + "extra": {"n": 1}, + }, + ] + assert IndexedTaskState.deserialize(state.serialize()).outlet_events == state.outlet_events + + def test_record_outlet_events_leaves_the_checkpoint_alone_when_nothing_was_recorded(self): + state = IndexedTaskState(status=TaskInstanceState.SUCCESS) + state.record_outlet_events(OutletEventAccessors()) + assert state.outlet_events is None + + def test_replay_outlet_events_lands_them_in_the_parents_accessors(self): + """A replayed checkpoint ends up as the live merge would leave it: one event per asset.""" + state = IndexedTaskState(status=TaskInstanceState.SUCCESS) + state.record_outlet_events(self._recorded_accessors()) + target = OutletEventAccessors() + target[Asset(name="a", uri="s3://bucket/a")].extra = {"rows": 1, "kept": True} + + IndexedTaskState.deserialize(state.serialize()).replay_outlet_events(target) + + asset = target[Asset(name="a", uri="s3://bucket/a")] + assert asset.extra == {"rows": 3, "kept": True} + assert asset.partition_keys == {"p1", "p2"} + alias_events = target[AssetAlias(name="alias")].asset_alias_events + assert [event.dest_asset_key for event in alias_events] == [ + AssetUniqueKey(name="b", uri="s3://bucket/b") + ] + assert len(list(target.items())) == 2 + + def test_replay_outlet_events_without_a_snapshot_does_nothing(self): + target = OutletEventAccessors() + IndexedTaskState(status=TaskInstanceState.SUCCESS).replay_outlet_events(target) + assert list(target.items()) == [] + + +class TestIndexedTaskInstance: + @pytest.mark.parametrize( + ("task_ids", "dag_id", "expected_key"), + [ + (None, None, "custom_key_2"), + ("the_task", None, "custom_key_2"), + ("the_task", "the_dag", "custom_key_2"), + ("upstream", None, "custom_key"), + (["the_task", "upstream"], None, "custom_key"), + ("the_task", "other_dag", "custom_key"), + ], + ids=["no_task", "own_task", "own_task_and_dag", "upstream", "several", "other_dag"], + ) + def test_xcom_pull_adds_the_index_for_its_own_xcoms_only( + self, make_indexed_ti, task_ids, dag_id, expected_key + ): + """What the iteration pushed under ``_`` is read back under that name; upstream keys stay.""" + ti = make_indexed_ti(index=2, task_id="the_task", dag_id="the_dag") + + with mock.patch.object(RuntimeTaskInstance, "xcom_pull", autospec=True) as pull: + ti.xcom_pull(task_ids=task_ids, dag_id=dag_id, key="custom_key") + + assert pull.call_args.kwargs["key"] == expected_key + assert pull.call_args.kwargs["task_ids"] == task_ids + + @pytest.mark.asyncio + async def test_axcom_pull_adds_the_index_for_its_own_xcoms_only(self, make_indexed_ti): + ti = make_indexed_ti(index=2, task_id="the_task", dag_id="the_dag") + + with mock.patch.object(RuntimeTaskInstance, "axcom_pull", autospec=True) as pull: + await ti.axcom_pull(key="custom_key") + await ti.axcom_pull(task_ids="upstream", key="custom_key") + + assert [call.kwargs["key"] for call in pull.call_args_list] == ["custom_key_2", "custom_key"] + + @pytest.mark.parametrize( + ("index", "key", "value", "expected_key"), + [ + (3, "result", "ok", "result_3"), + (2, BaseXCom.XCOM_RETURN_KEY, "value1", f"{BaseXCom.XCOM_RETURN_KEY}_2"), + (1, "custom_key", "value2", "custom_key_1"), + ], + ids=["delegates_with_index_suffix", "default_key", "custom_key"], + ) + def test_xcom_push_suffix(self, make_indexed_ti, index, key, value, expected_key): + """xcom_push appends the sub-task's index suffix to the XCom key.""" + ti = make_indexed_ti(index=index) + + with mock.patch("airflow.sdk.execution_time.task_runner._xcom_push", autospec=True) as mock_push: + ti.xcom_push(key=key, value=value) + + mock_push.assert_called_once_with(ti, expected_key, value) + + def test_properties(self, make_indexed_ti): + ti = make_indexed_ti(index=7, try_number=4, is_async=True, do_xcom_push=False) + + assert ti.is_async is True + assert ti.do_xcom_push is False + + @staticmethod + def _parent_context_and_operator(): + parent_operator = mock.create_autospec(BaseOperator, instance=True) + # The date the task is scheduled from, which is not when this task instance started. + parent_operator.start_date = timezone.datetime(2020, 1, 1) + parent = RuntimeTaskInstance.model_construct( + id=uuid7(), + task_id="iterated", + dag_id="dag", + run_id="run_1", + map_index=3, + try_number=2, + # What a manual clear leaves behind: the budget raised past the operator's retries. + max_tries=7, + start_date=timezone.datetime(2024, 12, 3, 9, 55, 0), + ) + context = {"ti": parent, "task": parent_operator, "task_state_store": parent.task_state_store} + operator = mock.create_autospec(BaseOperator, instance=True) + operator.task_id = "iterated" + operator.dag_id = "dag" + operator.retries = 5 + return context, operator + + def test_create_indexed_task_shares_the_parent_identity(self): + context, operator = self._parent_context_and_operator() + + ti = IndexedTaskInstance.create_indexed_task(context=context, index=4, operator=operator) + + parent = context["ti"] + assert (ti.id, ti.run_id, ti.map_index, ti.try_number) == (parent.id, "run_1", 3, 2) + # The retry budget is the parent's, not the operator's retries (5). + assert (ti.index, ti.max_tries, ti.task) == (4, 7, operator) + assert ti.start_date == parent.start_date + assert ti.state == TaskInstanceState.SCHEDULED.value + assert ti.is_mapped is True + + def test_create_indexed_task_carries_the_parents_context_from_the_server(self, create_runtime_ti): + """ + An iteration sees the dag run the task instance runs in: its logical date, and a template + context with ``dag_run`` and ``ds``, which ``get_previous_ti()`` and lineage macros read. + """ + parent = create_runtime_ti(task=BaseOperator(task_id="iterated")) + context = parent.get_template_context() + + ti = IndexedTaskInstance.create_indexed_task(context=context, index=0, operator=parent.task) + + assert ti._ti_context_from_server is parent._ti_context_from_server + assert ti.logical_date is not None + assert ti.logical_date == parent.logical_date + template_context = ti.get_template_context() + assert template_context["dag_run"] == context["dag_run"] + assert template_context["ds"] == context["ds"] + + def test_create_indexed_task_takes_the_parents_task_id(self): + """An iteration's XComs and state belong to the task instance running it, whatever the operator says.""" + context, operator = self._parent_context_and_operator() + operator.task_id = "group.iterated" + + ti = IndexedTaskInstance.create_indexed_task(context=context, index=4, operator=operator) + + assert ti.task_id == context["ti"].task_id == "iterated" + + def test_create_indexed_task_sees_the_parents_state_store_through_its_index(self): + """User code in an iteration gets the parent's store with suffixed keys; the checkpoints do not.""" + context, operator = self._parent_context_and_operator() + parent_store = context["ti"].task_state_store + + ti = IndexedTaskInstance.create_indexed_task(context=context, index=4, operator=operator) + + assert ti.parent_task_state_store is parent_store + assert ti.task_state_store == IndexedTaskStateStoreAccessor(parent_store, index=4) + assert ti.task_state_store is ti.task_state_store # cached, so the context and ti agree + + @pytest.mark.asyncio + async def test_pushes_other_than_the_return_value_are_recorded_by_their_bare_key(self): + """Sync and async pushes are recorded in memory, unsuffixed; the return value has its own slot.""" + context, operator = self._parent_context_and_operator() + ti = IndexedTaskInstance.create_indexed_task(context=context, index=4, operator=operator) + + with ( + mock.patch("airflow.sdk.execution_time.task_runner._xcom_push", autospec=True) as push, + mock.patch("airflow.sdk.execution_time.task_runner._axcom_push", autospec=True) as apush, + ): + ti.xcom_push(key="sync_key", value=1) + await ti.axcom_push(key="async_key", value=2) + ti.xcom_push(key=BaseXCom.XCOM_RETURN_KEY, value="result") + ti.xcom_push(key="sync_key", value=3) + + assert ti.pushed_xcoms == {"sync_key": 3, "async_key": 2} + assert [call.args[1] for call in push.call_args_list] == [ + "sync_key_4", + f"{BaseXCom.XCOM_RETURN_KEY}_4", + "sync_key_4", + ] + assert [call.args[1] for call in apush.call_args_list] == ["async_key_4"] + + def test_each_indexed_task_instance_records_its_own_pushes(self): + context, operator = self._parent_context_and_operator() + first = IndexedTaskInstance.create_indexed_task(context=context, index=0, operator=operator) + second = IndexedTaskInstance.create_indexed_task(context=context, index=1, operator=operator) + + with mock.patch("airflow.sdk.execution_time.task_runner._xcom_push", autospec=True): + first.xcom_push(key="foo", value="first") + + assert first.pushed_xcoms == {"foo": "first"} + assert second.pushed_xcoms == {} + + @pytest.mark.asyncio + async def test_checkpoints_go_to_the_parents_store_unsuffixed(self): + context, operator = self._parent_context_and_operator() + parent_store = context["ti"].task_state_store + ti = IndexedTaskInstance.create_indexed_task(context=context, index=4, operator=operator) + + with ( + mock.patch.object(parent_store, "aset", new_callable=mock.AsyncMock) as aset, + mock.patch.object(parent_store, "aget", new_callable=mock.AsyncMock, return_value=None) as aget, + ): + await ti.aset_state(IndexedTaskState(status=TaskInstanceState.SUCCESS)) + assert await ti.aget_state() is None + + aset.assert_awaited_once_with("_iterable_4", {"status": "success"}) + aget.assert_awaited_once_with("_iterable_4") + + def test_create_indexed_task_rejects_negative_index(self): + context, operator = self._parent_context_and_operator() + + with pytest.raises(ValueError, match="requires index >= 0, got -1"): + IndexedTaskInstance.create_indexed_task(context=context, index=-1, operator=operator) + + +class TestIndexedTaskRunner: + def test_dag_id_property(self, make_indexed_ti): + ti = make_indexed_ti(dag_id="my_dag") + executor = IndexedTaskRunner(task_instance=ti) + assert executor.dag_id == "my_dag" + + def test_task_id_property(self, make_indexed_ti): + ti = make_indexed_ti(task_id="my_task") + executor = IndexedTaskRunner(task_instance=ti) + assert executor.task_id == "my_task" + + def test_task_index(self, make_indexed_ti): + ti = make_indexed_ti(index=3) + executor = IndexedTaskRunner(task_instance=ti) + assert executor.task_index == ti.index + assert executor.task_index == 3 + + def test_operator_property(self, make_indexed_ti): + ti = make_indexed_ti() + executor = IndexedTaskRunner(task_instance=ti) + assert executor.operator is ti.task + + def test_is_async_property_sync(self, make_indexed_ti): + ti = make_indexed_ti(is_async=False) + executor = IndexedTaskRunner(task_instance=ti) + assert executor.is_async is False + + def test_is_async_property_async(self, make_indexed_ti): + ti = make_indexed_ti(is_async=True) + executor = IndexedTaskRunner(task_instance=ti) + assert executor.is_async is True + + def test_context_for_swaps_the_keys_the_indexed_task_owns(self, make_indexed_ti): + """ + ``ti``, ``task_instance``, ``task`` and ``task_state_store`` become the indexed task's on a + copy of the parent's context; the rest, outlet events included, comes over as it is. + """ + ti = make_indexed_ti(index=3) + ti.parent_task_state_store = mock.MagicMock(name="parent_store") + parent_ti = mock.MagicMock(name="parent_ti") + parent_events = mock.MagicMock(name="parent_events") + parent = { + "ti": parent_ti, + "task_instance": parent_ti, + "task": parent_ti.task, + "task_state_store": parent_ti.task_state_store, + "outlet_events": parent_events, + "params": {"p": 1}, + } + + context = ti.context_for(parent) + + assert context["ti"] is ti + assert context["task_instance"] is ti + assert context["task"] is ti.task + assert context["task_state_store"] is ti.task_state_store + assert context["outlet_events"] is parent_events + assert context["params"] == {"p": 1} + assert context["params"] is not parent["params"] + assert parent["ti"] is parent_ti + + def test_context_for_swaps_the_outlet_events_when_given(self, make_indexed_ti): + ti = make_indexed_ti(index=3) + ti.parent_task_state_store = mock.MagicMock(name="parent_store") + own_events = OutletEventAccessors() + parent = { + "ti": mock.MagicMock(name="parent_ti"), + "outlet_events": mock.MagicMock(name="parent_events"), + } + + context = ti.context_for(parent, outlet_events=own_events) + + assert context["outlet_events"] is own_events + + def test_indexed_context_is_the_parents_seen_from_the_indexed_task(self, make_indexed_ti): + """ + Inside the block the indexed task has its own ti, state store view and outlet events on a + copy of the parent's context, and that copy is the current context. + """ + ti = make_indexed_ti(index=3) + ti.parent_task_state_store = mock.MagicMock(name="parent_store") + parent_ti = mock.MagicMock(name="parent_ti") + parent = { + "ti": parent_ti, + "task_instance": parent_ti, + "task": parent_ti.task, + "task_state_store": parent_ti.task_state_store, + "outlet_events": mock.MagicMock(name="parent_events"), + "inlet_events": mock.MagicMock(name="inlet_events"), + "dag_run": mock.MagicMock(name="dag_run"), + "params": {"p": 1}, + } + runner = IndexedTaskRunner(task_instance=ti) + + with runner.indexed_context(parent) as indexed_context: + assert get_current_context() is indexed_context + assert indexed_context["ti"] is ti + assert indexed_context["task_instance"] is ti + assert indexed_context["task"] is ti.task + assert indexed_context["task_state_store"] is ti.task_state_store + assert indexed_context["outlet_events"] is runner.outlet_events + assert indexed_context["params"] == {"p": 1} + # a copy, so concurrent indexed tasks cannot leak into each other + assert indexed_context["params"] is not parent["params"] + + assert runner._context is indexed_context # remembered for the state-change callbacks + assert parent["ti"] is parent_ti + assert parent["task"] is parent_ti.task + assert parent["outlet_events"] is not runner.outlet_events + + def test_outlet_events_can_be_handed_in(self, make_indexed_ti): + events = mock.MagicMock(name="events") + executor = IndexedTaskRunner(task_instance=make_indexed_ti(), outlet_events=events) + assert executor.outlet_events is events + + def test_merge_outlet_events_into_folds_them_into_one_event_per_asset(self, make_indexed_ti): + """Siblings emitting to one asset share its event: the last extra wins, the rest accumulates.""" + parent = OutletEventAccessors() + asset = Asset(name="a", uri="s3://bucket/a") + first = IndexedTaskRunner(task_instance=make_indexed_ti(index=0)) + first.outlet_events[asset].extra = {"rows": 1, "from": "first"} + first.outlet_events[asset].add_partitions(["p1"]) + second = IndexedTaskRunner(task_instance=make_indexed_ti(index=1)) + second.outlet_events[asset].extra = {"rows": 2} + second.outlet_events[asset].add_partitions(["p2"]) + second.outlet_events[AssetAlias(name="alias")].asset_alias_events.append( + AssetAliasEvent( + source_alias_name="alias", + dest_asset_key=AssetUniqueKey(name="b", uri="s3://bucket/b"), + dest_asset_extra={}, + extra={}, + ) + ) + + first.merge_outlet_events_into(parent) + second.merge_outlet_events_into(parent) + + assert parent[asset].extra == {"rows": 2, "from": "first"} + assert parent[asset].partition_keys == {"p1", "p2"} + assert len(parent[AssetAlias(name="alias")].asset_alias_events) == 1 + assert len(list(parent.items())) == 2 + + def test_enter_sets_start_time(self, make_indexed_ti): + ti = make_indexed_ti() + executor = IndexedTaskRunner(task_instance=ti) + assert executor._start_time is None + executor.__enter__() + assert executor._start_time is not None + + def test_enter_returns_self(self, make_indexed_ti): + ti = make_indexed_ti() + executor = IndexedTaskRunner(task_instance=ti) + with executor as ctx: + assert ctx is executor + + def test_in_flight_registers_the_operator_while_its_code_runs(self, make_indexed_ti): + """run()/arun() enter in_flight() in the thread or coroutine running the operator, not __enter__.""" + ti = make_indexed_ti() + state = IterationState() + runner = IndexedTaskRunner(task_instance=ti, register=state) + + with runner: + assert ti.task not in state + with runner.in_flight(): + assert ti.task in state + assert ti.task not in state + + def test_in_flight_unregisters_on_failure(self, make_indexed_ti): + """The operator leaves the register even when it raises, so on_kill() never reaches a stopped one.""" + ti = make_indexed_ti() + state = IterationState() + runner = IndexedTaskRunner(task_instance=ti, register=state) + + with pytest.raises(ValueError, match="boom"): + with runner.in_flight(): + raise ValueError("boom") + + assert ti.task not in state + ti.task.on_kill.assert_not_called() + + def test_in_flight_leaves_the_operator_the_parent_timeout_strikes_registered(self, make_indexed_ti): + """ + The timeout lands on the loop thread; the operator's on_kill must not run there, so the + operator stays registered for the parent's kill off the loop once the timeout has unwound. + """ + ti = make_indexed_ti() + state = IterationState() + runner = IndexedTaskRunner(task_instance=ti, register=state) + + with pytest.raises(AirflowTaskTimeout): + with runner.in_flight(): + raise AirflowTaskTimeout("the task ran out of time") + + ti.task.on_kill.assert_not_called() + assert ti.task in state + + def test_in_flight_keeps_operators_that_compare_equal_apart(self, make_indexed_ti): + """Sub-operators of one iterated task compare equal; each is registered on its own.""" + first, second = make_indexed_ti(index=0), make_indexed_ti(index=1) + first.task, second.task = BaseOperator(task_id="same"), BaseOperator(task_id="same") + assert first.task == second.task + state = IterationState() + + with IndexedTaskRunner(task_instance=first, register=state).in_flight(): + with IndexedTaskRunner(task_instance=second, register=state).in_flight(): + assert len(state.take_in_flight()) == 2 + + def test_the_register_is_optional(self, make_indexed_ti): + """When no register is supplied (the default), in_flight() must not raise.""" + ti = make_indexed_ti() + runner = IndexedTaskRunner(task_instance=ti) + with runner, runner.in_flight(): + pass # should not raise + + def test_exit_success_leaves_the_state_to_report_success(self, make_indexed_ti): + """__exit__ without an exception changes no state; report_success() marks the task instance SUCCESS.""" + ti = make_indexed_ti() + state_before = ti.state + runner = IndexedTaskRunner(task_instance=ti) + with runner: + pass # no exception + assert ti.state == state_before + + runner.report_success() + + assert ti.state == TaskInstanceState.SUCCESS + + def test_exit_with_task_deferred_reraises(self, make_indexed_ti): + """TaskDeferred must propagate unchanged through __exit__.""" + ti = make_indexed_ti() + trigger = mock.Mock() + deferred = TaskDeferred(trigger=trigger, method_name="resume") + + with pytest.raises(TaskDeferred): + with IndexedTaskRunner(task_instance=ti): + raise deferred + + @staticmethod + def _parent_context(task): + """A parent context with what clone_context needs and a parent-side state store.""" + context = mock_context(task) + context["inlet_events"] = mock.MagicMock(name="inlet_events") + context["dag_run"] = mock.MagicMock(name="dag_run") + context["task_state_store"] = mock.MagicMock(name="parent_store") + return context + + def test_run_delegates_to_execute_task(self, make_indexed_ti): + """run() must call _execute_task with the sub-task's view of the given context.""" + ti = make_indexed_ti() + ti.parent_task_state_store = mock.MagicMock(name="parent_store") + task = BaseOperator(task_id="test_task") + get_inline_dag("test_dag", task) + context = self._parent_context(task) + executor = IndexedTaskRunner(task_instance=ti) + + with mock.patch( + "airflow.sdk.execution_time.task_runner._execute_task", + autospec=True, + return_value="result", + ) as mock_execute: + result = executor.run(context) + + mock_execute.assert_called_once() + indexed_context, passed_ti, log = mock_execute.call_args.args + assert (passed_ti, log) == (ti, executor.log) + assert indexed_context["ti"] is ti + assert indexed_context["task_state_store"] is ti.task_state_store + assert indexed_context["outlet_events"] is executor.outlet_events + assert context["ti"] is not ti # the parent's context is untouched + assert result == "result" + + @pytest.mark.asyncio + async def test_arun_delegates_to_execute_async_task(self, make_indexed_ti): + """arun() must call _execute_async_task with the sub-task's view of the given context.""" + ti = make_indexed_ti(is_async=True) + ti.parent_task_state_store = mock.MagicMock(name="parent_store") + task = BaseOperator(task_id="test_task") + get_inline_dag("test_dag", task) + context = self._parent_context(task) + executor = IndexedTaskRunner(task_instance=ti) + + with mock.patch( + "airflow.sdk.execution_time.task_runner._execute_async_task", + new=mock.AsyncMock(return_value="async_result"), + ) as mock_async_execute: + result = await executor.arun(context) + + mock_async_execute.assert_called_once() + indexed_context, passed_ti, log = mock_async_execute.call_args.args + assert (passed_ti, log) == (ti, executor.log) + assert indexed_context["ti"] is ti + assert indexed_context["task_state_store"] is ti.task_state_store + assert result == "async_result" + + def test_exit_success_reports_nothing_until_report_success(self, make_indexed_ti): + """ + A success is reported once its checkpoint is written, by report_success(), not by the exit. + + The callback then speaks for work a retry will not run again; it gets the state and an + ``end_date`` on the indexed task instance, as a plain task's callback does. + """ + fired: list[str] = [] + ti = make_indexed_ti() + task = BaseOperator( + task_id="cb_task", + on_success_callback=lambda ctx: fired.append(("success", ti.state, ti.end_date)), + ) + get_inline_dag("cb_dag", task) + ti.task = task + state_before = ti.state + context = mock_context(task) + executor = IndexedTaskRunner(task_instance=ti) + executor._context = context + + with executor: + pass + + assert fired == [] + assert ti.state == state_before + + executor.report_success() + + assert fired == [("success", TaskInstanceState.SUCCESS, ti.end_date)] + assert ti.end_date is not None + + @pytest.mark.parametrize( + "exception", + [ + RuntimeError("transient"), + AirflowFailException("do not retry"), + AirflowTaskTimeout("the task ran out of time"), + SystemExit(1), + ], + ids=["error", "fail", "timeout", "system-exit"], + ) + def test_exit_notes_a_failure_without_deciding_it(self, make_indexed_ti, exception): + """Whether a failure is retried is the task's fate: __exit__ only notes it, no state, no callback.""" + fired: list[str] = [] + ti = make_indexed_ti(try_number=1, max_tries=3) + task = BaseOperator( + task_id="cb_task", + on_failure_callback=lambda ctx: fired.append("failure"), + on_retry_callback=lambda ctx: fired.append("retry"), + ) + get_inline_dag("cb_dag", task) + ti.task = task + state_before = ti.state + runner = IndexedTaskRunner(task_instance=ti) + runner._context = mock_context(task) + + with pytest.raises(type(exception)): + with runner: + raise exception + + assert runner.failure is exception + assert fired == [] + assert ti.state == state_before + + @pytest.mark.parametrize( + ("task_will_retry", "fired_callback", "state"), + [ + pytest.param(True, "retry", TaskInstanceState.UP_FOR_RETRY, id="retried"), + pytest.param(False, "failure", TaskInstanceState.FAILED, id="failed"), + ], + ) + def test_report_failure_follows_the_tasks_fate( + self, make_indexed_ti, task_will_retry, fired_callback, state + ): + fired: list[str] = [] + ti = make_indexed_ti(try_number=1, max_tries=3) + task = BaseOperator( + task_id="cb_task", + on_failure_callback=lambda ctx: fired.append("failure"), + on_retry_callback=lambda ctx: fired.append("retry"), + ) + get_inline_dag("cb_dag", task) + ti.task = task + runner = IndexedTaskRunner(task_instance=ti) + runner._context = mock_context(task) + with pytest.raises(RuntimeError): + with runner: + raise RuntimeError("boom") + + runner.report_failure(task_will_retry=task_will_retry) + + assert fired == [fired_callback] + assert ti.state == state + + @pytest.mark.parametrize( + ("try_number", "max_tries", "eligible"), + [ + (1, 0, False), # no retries: the only attempt fails + (1, 1, True), # one retry: the first attempt is retried + (2, 1, False), # one retry: the second attempt fails + (3, 3, True), # three retries: the third attempt is still retried + (4, 3, False), # three retries: the fourth attempt fails + (2, 3, True), # cleared after one attempt with two retries: budget raised to 3 + ], + ) + def test_is_eligible_to_retry_follows_the_server_rule( + self, make_indexed_ti, try_number, max_tries, eligible + ): + """The attempt numbers the parent really has (the first is 1) and ``try_number <= max_tries``.""" + assert make_indexed_ti(try_number=try_number, max_tries=max_tries).is_eligible_to_retry is eligible + + def test_exit_skip_reports_nothing_until_report_skip(self, make_indexed_ti): + """ + A skipped iteration is neither a failure nor a retry, whatever budget is left. + + The exit lets the skip through untouched; report_skip(), once the SKIPPED checkpoint is + written, sets the state and ``end_date`` and fires the skipped callback only. + """ + fired: list[str] = [] + ti = make_indexed_ti(try_number=1, max_tries=3) + task = BaseOperator( + task_id="cb_task", + on_skipped_callback=lambda ctx: fired.append("skipped"), + on_failure_callback=lambda ctx: fired.append("failure"), + on_retry_callback=lambda ctx: fired.append("retry"), + ) + get_inline_dag("cb_dag", task) + ti.task = task + state_before = ti.state + executor = IndexedTaskRunner(task_instance=ti) + executor._context = mock_context(task) + + with pytest.raises(AirflowSkipException): + with executor: + raise AirflowSkipException("nothing to do") + + assert fired == [] + assert ti.state == state_before + assert executor.failure is None + + executor.report_skip() + + assert fired == ["skipped"] + assert ti.state == TaskInstanceState.SKIPPED + assert ti.end_date is not None + + def test_exit_cancelled_iteration_gets_no_state_and_no_callback(self, make_indexed_ti): + """Cancelled because the task is stopping: the iteration neither failed nor will be retried.""" + fired: list[str] = [] + ti = make_indexed_ti(try_number=1, max_tries=3) + task = BaseOperator( + task_id="cb_task", + on_failure_callback=lambda ctx: fired.append("failure"), + on_retry_callback=lambda ctx: fired.append("retry"), + ) + get_inline_dag("cb_dag", task) + ti.task = task + state_before = ti.state + executor = IndexedTaskRunner(task_instance=ti) + executor._context = mock_context(task) + + with pytest.raises(asyncio.CancelledError): + with executor: + raise asyncio.CancelledError() + + assert fired == [] + assert ti.state == state_before + + def test_report_without_a_context_skips_callbacks(self, make_indexed_ti): + """When _context is not set (the runner never entered its context), callbacks must not fire.""" + fired: list[str] = [] + ti = make_indexed_ti() + task = BaseOperator( + task_id="cb_task", + on_success_callback=lambda ctx: fired.append("success"), + on_skipped_callback=lambda ctx: fired.append("skipped"), + ) + get_inline_dag("cb_dag", task) + ti.task = task + executor = IndexedTaskRunner(task_instance=ti) + # _context intentionally left as None + + with executor: + pass + executor.report_success() + executor.report_skip() + + assert fired == [] + + class TestSerializeOutletEvents: """Tests for the wire format produced by ``_serialize_outlet_events``.""" @@ -2810,6 +3651,86 @@ def test_lazy_loading_not_triggered_until_accessed(self, create_runtime_ti, mock # Now the lazy attribute should trigger the call mock_supervisor_comms.send.assert_called_once() + def test_logical_date_returns_none_without_ti_context_from_server(self, mocked_parse): + """Test that logical_date returns None when _ti_context_from_server is not set.""" + task = BaseOperator(task_id="hello") + dag_id = "basic_task" + + get_inline_dag(dag_id=dag_id, task=task) + + ti_id = uuid7() + ti = TaskInstance( + id=ti_id, + task_id=task.task_id, + dag_id=dag_id, + run_id="test_run", + try_number=1, + dag_version_id=uuid7(), + ) + start_date = timezone.datetime(2025, 1, 1) + + runtime_ti = RuntimeTaskInstance.model_construct( + **ti.model_dump(exclude_unset=True), + task=task, + _ti_context_from_server=None, + start_date=start_date, + ) + + assert runtime_ti.logical_date is None + + def test_logical_date_returns_dag_run_logical_date(self, create_runtime_ti): + """Test that logical_date returns the dag run's logical_date when _ti_context_from_server is set.""" + task = BaseOperator(task_id="hello") + runtime_ti = create_runtime_ti(task=task, dag_id="basic_task") + + dag_run = runtime_ti._ti_context_from_server.dag_run + + assert runtime_ti.logical_date == dag_run.logical_date + assert runtime_ti.logical_date == timezone.datetime(2024, 12, 1, 1, 0, 0) + + def test_task_state_store_is_cached(self, create_runtime_ti): + """Repeated access must return the same instance, not rebuild a new accessor each time.""" + task = BaseOperator(task_id="hello") + runtime_ti = create_runtime_ti(task=task, dag_id="basic_task") + + first = runtime_ti.task_state_store + second = runtime_ti.task_state_store + + assert first is second + + def test_task_state_store_used_by_template_context_is_the_cached_instance(self, create_runtime_ti): + """``get_template_context()`` must wire in the same cached accessor, not a fresh one.""" + task = BaseOperator(task_id="hello") + runtime_ti = create_runtime_ti(task=task, dag_id="basic_task") + + context = runtime_ti.get_template_context() + + assert context["task_state_store"] is runtime_ti.task_state_store + + @pytest.mark.parametrize( + ("map_index", "expected_scope_map_index"), + [ + pytest.param(None, -1, id="explicit-none-map-index-falls-back-to-minus-one"), + pytest.param(0, 0, id="mapped-task-index-zero"), + pytest.param(3, 3, id="mapped-task-index-three"), + ], + ) + def test_task_state_store_scope_reflects_map_index( + self, create_runtime_ti, map_index, expected_scope_map_index + ): + """The scope used to namespace task-state-store keys must match the TI's own map_index, + falling back to -1 when map_index is None (e.g. an unmapped task).""" + task = BaseOperator(task_id="hello") + runtime_ti = create_runtime_ti(task=task, dag_id="basic_task", map_index=map_index) + assert runtime_ti.map_index == map_index + + scope = runtime_ti.task_state_store._scope + + assert scope.map_index == expected_scope_map_index + assert scope.dag_id == runtime_ti.dag_id + assert scope.run_id == runtime_ti.run_id + assert scope.task_id == runtime_ti.task_id + def test_get_connection_from_context(self, create_runtime_ti, mock_supervisor_comms): """Test that the connection is fetched from the API server via the Supervisor lazily when accessed""" diff --git a/task-sdk/tests/task_sdk/serde/test_serde.py b/task-sdk/tests/task_sdk/serde/test_serde.py index 37d93a6e89578..4838810593746 100644 --- a/task-sdk/tests/task_sdk/serde/test_serde.py +++ b/task-sdk/tests/task_sdk/serde/test_serde.py @@ -30,6 +30,7 @@ from pydantic import BaseModel, create_model from airflow._shared.module_loading import import_string, iter_namespace, qualname +from airflow.sdk.bases.xcom import XComIterable from airflow.sdk.definitions.asset import Asset from airflow.sdk.serde import ( CLASSNAME, @@ -335,6 +336,31 @@ def test_serder_dataclass(self): d = deserialize(e) assert i.x == getattr(d, "x", None) + @pytest.mark.parametrize( + "instance", + [ + pytest.param( + lambda: XComIterable(task_id="t", dag_id="d", run_id="r", map_index=0, length=3), + id="XComIterable", + ), + ], + ) + def test_serder_xcom_iterable_round_trip(self, instance): + """XComIterable round-trips through serde without explicit registration.""" + obj = instance() + serialized = serialize(obj) + assert serialized[CLASSNAME] == qualname(obj) + assert serialized[DATA] + + restored = deserialize(serialized) + assert type(restored) is type(obj) + assert isinstance(restored, XComIterable) + assert restored.task_id == obj.task_id + assert restored.dag_id == obj.dag_id + assert restored.run_id == obj.run_id + assert restored.map_index == obj.map_index + assert restored.length == obj.length + @conf_vars( { ("core", "allowed_deserialization_classes"): "airflow.*",