dabla commented on code in PR #62922: URL: https://github.com/apache/airflow/pull/62922#discussion_r4141183899
########## task-sdk/src/airflow/sdk/definitions/iterableoperator.py: ########## @@ -0,0 +1,840 @@ +# +# 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 warnings +from collections.abc import AsyncIterable, AsyncIterator, 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 BaseAsyncOperator, BaseOperator, event_loop +from airflow.sdk.bases.xcom import XComIterable +from airflow.sdk.definitions.asset import Asset, AssetAlias, AssetAliasEvent, AssetUniqueKey +from airflow.sdk.definitions.xcom_arg import XComArg +from airflow.sdk.exceptions import ( + AirflowFailException, + AirflowRescheduleException, + AirflowSkipException, + AirflowTaskTimeout, + DagRunTriggerException, + DownstreamTasksSkipped, + TaskDeferred, +) +from airflow.sdk.execution_time.comms import DeadlockImminentError +from airflow.sdk.execution_time.context import OutletEventAccessors, 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 OutletEventAccessorsProtocol + + +# 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} +) + + +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 items. An index then no longer means the + same work, and replaying its result would hand downstream a value computed from the old item. + 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 _serialize_outlet_events(accessors: OutletEventAccessors) -> list[dict[str, Any]]: + """ + Snapshot the outlet asset events one sub-task recorded into a JSON-safe list. + + Persisted on the sub-task's checkpoint so a later attempt can replay them via + ``_replay_outlet_events`` when the sub-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, + } + ) + return events + + +def _merge_outlet_events(target: OutletEventAccessorsProtocol, source: OutletEventAccessors) -> None: + """ + Merge every outlet asset event recorded in ``source`` into ``target``. + + Used both to fold a sub-task's isolated accessor into the IterableOperator's shared + ``context["outlet_events"]`` right after it succeeds, and to replay a checkpointed + snapshot (via ``_replay_outlet_events``) for a sub-task skipped on retry. + """ + for asset_or_alias, accessor in source.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 _replay_outlet_events(target: OutletEventAccessorsProtocol, events: list[dict[str, Any]]) -> None: + """ + Re-populate ``target`` with events a sub-task recorded on a previous attempt. + + 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. 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. + """ + replayed = OutletEventAccessors() + for event in events: + if event["kind"] == "asset": + accessor = replayed[Asset(name=event["name"], uri=event["uri"])] + accessor.extra.update(event["extra"]) + if event["partition_keys"]: + accessor.add_partitions(event["partition_keys"]) + else: + accessor = replayed[AssetAlias(name=event["source_alias_name"])] + accessor.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"], + ) + ) + _merge_outlet_events(target, replayed) + + +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 item + 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 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`). + + :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 item 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_<index>`` 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 + (:class:`~airflow.sdk.exceptions.DownstreamTasksSkipped`, raised e.g. by + ``ShortCircuitOperator``) 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. Either exception raised by a + sub-task 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 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``. 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:: + **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``. ``Variable`` has no async equivalent yet. + + .. warning:: + **``execution_timeout`` caps the whole iteration; per-sub-task enforcement is async-only.** + + 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 items in one task instance, this is the per-instance limit of ``.expand()`` applied to + the whole iteration rather than to each item. + + Per item, only async sub-tasks (instances of :class:`~airflow.sdk.bases.operator.BaseAsyncOperator`) + are additionally limited, via ``asyncio.wait_for``. Sync sub-tasks run in worker threads and rely + on :class:`~airflow.sdk.execution_time.timeout.TimeoutPosix`, which requires ``signal.SIGALRM`` and + only works in the main thread, so no per-item limit applies to them. Use + :class:`~airflow.sdk.bases.operator.BaseAsyncOperator` if per-sub-task time limits are required. + """ + + _operator: MappedOperator + expand_input: ExpandInput + partial_kwargs: dict[str, Any] + shallow_copy_attrs: Sequence[str] = ( + "_operator", + "expand_input", + "partial_kwargs", + "_log", + "_active_sub_operators", + "_active_sub_operators_lock", + "_resolved", + ) + + 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") + + super().__init__( + **{ + **kwargs, + "task_id": operator.task_id, Review Comment: Reproduced, for `@task` and classic operators alike: inside `TaskGroup("tg")` the iterated task was registered as `tg.tg.f`. 12cc0077cb passes the bare id to `BaseOperator.__init__` and makes `create_indexed_task` take `parent.task_id`, as you suggested. There are tests for one group, nested groups and `prefix_group_id=False`, plus a DB test where an iterated task inside a group feeds a downstream `.expand()`. Drafted-by: Claude Opus 5.5; reviewed by @dabla before posting -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
