dabla commented on code in PR #62922:
URL: https://github.com/apache/airflow/pull/62922#discussion_r4094332652


##########
task-sdk/src/airflow/sdk/definitions/iterableoperator.py:
##########
@@ -0,0 +1,487 @@
+#
+# 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 copy
+import os
+import warnings
+from collections.abc import Iterable, Mapping, Sequence
+from itertools import repeat
+from typing import TYPE_CHECKING, Any
+
+try:
+    # Python 3.11+
+    BaseExceptionGroup
+except NameError:
+    from exceptiongroup import BaseExceptionGroup
+
+from airflow.sdk import BaseXCom, TaskInstanceState
+from airflow.sdk.bases.operator import BaseAsyncOperator, BaseOperator, 
event_loop
+from airflow.sdk.definitions._internal.expandinput import BatchedExpandInput
+from airflow.sdk.definitions.context import clone_context
+from airflow.sdk.definitions.mappedoperator import MappedOperator
+from airflow.sdk.definitions.xcom_arg import MapXComArg, XComArg  # noqa: F401
+from airflow.sdk.exceptions import (
+    AirflowFailException,
+    TaskDeferred,
+)
+from airflow.sdk.execution_time.executor import AsyncAwareExecutor, 
TaskExecutor
+from airflow.sdk.execution_time.task_runner import IndexedTaskInstance
+
+if TYPE_CHECKING:
+    import jinja2
+
+    from airflow.sdk.bases.xcom import XComIterable
+    from airflow.sdk.definitions._internal.expandinput import ExpandInput
+    from airflow.sdk.definitions.context import Context
+
+_ITERABLE_CHECKPOINT_KEY_PREFIX = "_iterable_task_"
+
+
+class IterableOperator(BaseOperator):
+    """
+    Operator used for Task Iteration (TI) 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``.
+
+    **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 pending/failed 
sub-tasks are recreated
+    with their accumulated ``try_number`` so that retries are not wasted.
+
+    :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 (or batches) to iterate
+        over. Its ``iter_values(context)`` method is used to produce 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::
+        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 meaningfully supported: the exception is treated like any 
other sub-task failure and
+        counts towards the IterableOperator's own ``retries``, but the 
requested ``reschedule_date`` is not
+        honored — the worker is not released and the next attempt follows the 
IterableOperator's own
+        ``retry_delay`` instead of waiting until ``reschedule_date``.
+
+    .. warning::
+        **``execution_timeout`` is only enforced for async sub-tasks.**
+
+        Async sub-tasks (instances of 
:class:`~airflow.sdk.bases.operator.BaseAsyncOperator`) respect
+        ``execution_timeout`` 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. Because sync sub-tasks execute in a 
thread pool, ``SIGALRM`` cannot
+        be delivered to them, so their ``execution_timeout`` is silently 
ignored. 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",
+    )
+
+    def __init__(
+        self,
+        *,
+        operator: MappedOperator,
+        expand_input: ExpandInput,
+        **kwargs,
+    ):
+        super().__init__(
+            **{
+                **kwargs,
+                "task_id": operator.task_id,
+                "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,
+                "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,
+                "priority_weight": operator.priority_weight,
+                "queue": operator.queue,
+                "pool": operator.pool,
+                "pool_slots": operator.pool_slots,
+                "execution_timeout": None,
+                "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,
+                "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,
+            }
+        )
+        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}")
+        # Known v1 limitation: pool_slots is reserved once by the scheduler 
for this IterableOperator TI,
+        # but up to max_workers sub-tasks run concurrently inside it. 
Operators that set pool_slots > 1 to
+        # protect a shared resource (e.g. a DB connection pool) will be 
under-accounted — the pool sees one
+        # reservation while max_workers connections can be active 
simultaneously. A proper fix requires the
+        # scheduler to reserve pool_slots * max_workers slots, which needs 
scheduler-side changes.
+        self.max_workers = task_concurrency if task_concurrency is not None 
else (os.cpu_count() or 1)
+        if operator.execution_timeout and not 
issubclass(operator.operator_class, BaseAsyncOperator):
+            warnings.warn(
+                f"Operator {operator.task_id!r} has execution_timeout set, but 
sync operators run in "
+                "worker threads where TimeoutPosix (SIGALRM) cannot be 
delivered. "
+                "The execution_timeout will not be enforced for sync sub-tasks 
inside IterableOperator. "
+                "Use BaseAsyncOperator if per-sub-task time limits are 
required.",
+                UserWarning,
+                stacklevel=2,
+            )
+        XComArg.apply_upstream_relationship(self, self.expand_input.value)
+
+    @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 task_type(self) -> str:
+        return self._operator.__class__.__name__

Review Comment:
   Fixed in c8fb12d1d7 (`task_type`) and 93898fba9c (`operator_name`): both 
forward to the wrapped operator, so an `.iterate()` task shows its real class 
in the UI and API, matching what the `.batch()` path already did. Two tests pin 
it.
   
   ---
   Drafted-by: Claude Fable 5.1; reviewed by @dabla before posting



##########
task-sdk/src/airflow/sdk/definitions/iterableoperator.py:
##########
@@ -0,0 +1,487 @@
+#
+# 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 copy
+import os
+import warnings
+from collections.abc import Iterable, Mapping, Sequence
+from itertools import repeat
+from typing import TYPE_CHECKING, Any
+
+try:
+    # Python 3.11+
+    BaseExceptionGroup
+except NameError:
+    from exceptiongroup import BaseExceptionGroup
+
+from airflow.sdk import BaseXCom, TaskInstanceState
+from airflow.sdk.bases.operator import BaseAsyncOperator, BaseOperator, 
event_loop
+from airflow.sdk.definitions._internal.expandinput import BatchedExpandInput
+from airflow.sdk.definitions.context import clone_context
+from airflow.sdk.definitions.mappedoperator import MappedOperator
+from airflow.sdk.definitions.xcom_arg import MapXComArg, XComArg  # noqa: F401
+from airflow.sdk.exceptions import (
+    AirflowFailException,
+    TaskDeferred,
+)
+from airflow.sdk.execution_time.executor import AsyncAwareExecutor, 
TaskExecutor
+from airflow.sdk.execution_time.task_runner import IndexedTaskInstance
+
+if TYPE_CHECKING:
+    import jinja2
+
+    from airflow.sdk.bases.xcom import XComIterable
+    from airflow.sdk.definitions._internal.expandinput import ExpandInput
+    from airflow.sdk.definitions.context import Context
+
+_ITERABLE_CHECKPOINT_KEY_PREFIX = "_iterable_task_"
+
+
+class IterableOperator(BaseOperator):
+    """
+    Operator used for Task Iteration (TI) 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``.
+
+    **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 pending/failed 
sub-tasks are recreated
+    with their accumulated ``try_number`` so that retries are not wasted.
+
+    :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 (or batches) to iterate
+        over. Its ``iter_values(context)`` method is used to produce 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::
+        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 meaningfully supported: the exception is treated like any 
other sub-task failure and
+        counts towards the IterableOperator's own ``retries``, but the 
requested ``reschedule_date`` is not
+        honored — the worker is not released and the next attempt follows the 
IterableOperator's own
+        ``retry_delay`` instead of waiting until ``reschedule_date``.
+
+    .. warning::
+        **``execution_timeout`` is only enforced for async sub-tasks.**
+
+        Async sub-tasks (instances of 
:class:`~airflow.sdk.bases.operator.BaseAsyncOperator`) respect
+        ``execution_timeout`` 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. Because sync sub-tasks execute in a 
thread pool, ``SIGALRM`` cannot
+        be delivered to them, so their ``execution_timeout`` is silently 
ignored. 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",
+    )
+
+    def __init__(
+        self,
+        *,
+        operator: MappedOperator,
+        expand_input: ExpandInput,
+        **kwargs,
+    ):
+        super().__init__(

Review Comment:
   Fixed in 770263a15f. `params`, `weight_rule`, `retry_policy` and 
`do_xcom_push` are forwarded; `is_setup`, `is_teardown` and 
`on_failure_fail_dagrun` are applied from `partial_kwargs` the way `unmap` and 
`__attrs_post_init__` would; upstream edges are recorded for XComArgs in 
template-field partial kwargs, not only for `expand_input`; and the guard 
against `.iterate()` inside a mapped task group is back. Each has its own test.
   
   ---
   Drafted-by: Claude Fable 5.1; reviewed by @dabla before posting



##########
task-sdk/src/airflow/sdk/definitions/mappedoperator.py:
##########
@@ -213,68 +215,31 @@ def expand_kwargs(self, kwargs: 
OperatorExpandKwargsArgument, *, strict: bool =
             raise TypeError(f"expected XComArg or list[dict], not 
{type(kwargs).__name__}")
         return self._expand(ListOfDictsExpandInput(kwargs), strict=strict)
 
-    def _expand(self, expand_input: ExpandInput, *, strict: bool) -> 
MappedOperator:
-        from airflow.providers.standard.operators.empty import EmptyOperator
-        from airflow.sdk import BaseSensorOperator
-        from airflow.sdk.bases.skipmixin import SkipMixin
-
+    def _expand(
+        self,
+        expand_input: ExpandInput,
+        *,
+        strict: bool,
+        register_with_dag: bool = True,
+    ) -> MappedOperator:
         self._expand_called = True
-        ensure_xcomarg_return_value(expand_input.value)
-
-        partial_kwargs = self.kwargs.copy()
-        task_id = partial_kwargs.pop("task_id")
-        dag = partial_kwargs.pop("dag")
-        task_group = partial_kwargs.pop("task_group")
-        start_date = partial_kwargs.pop("start_date", None)
-        end_date = partial_kwargs.pop("end_date", None)
-        start_from_trigger = (
-            partial_kwargs["start_from_trigger"]
-            if "start_from_trigger" in partial_kwargs
-            else getattr(self.operator_class, "start_from_trigger", False)
-        )
-        start_trigger_args = (
-            partial_kwargs["start_trigger_args"]
-            if "start_trigger_args" in partial_kwargs
-            else getattr(self.operator_class, "start_trigger_args", None)
-        )
+        return self.batch(size=0)._expand(expand_input, strict=strict, 
register_with_dag=register_with_dag)
 
-        try:
-            operator_name = self.operator_class.custom_operator_name  # type: 
ignore
-        except AttributeError:
-            operator_name = self.operator_class.__name__
-
-        op = MappedOperator(
-            operator_class=self.operator_class,
-            expand_input=expand_input,
-            partial_kwargs=partial_kwargs,
-            task_id=task_id,
-            params=self.params,
-            operator_extra_links=self.operator_class.operator_extra_links,
-            template_ext=self.operator_class.template_ext,
-            template_fields=self.operator_class.template_fields,
-            
template_fields_renderers=self.operator_class.template_fields_renderers,
-            ui_color=self.operator_class.ui_color,
-            ui_fgcolor=self.operator_class.ui_fgcolor,
-            is_empty=issubclass(self.operator_class, EmptyOperator),
-            is_sensor=issubclass(self.operator_class, BaseSensorOperator),
-            can_skip_downstream=issubclass(self.operator_class, SkipMixin),
-            is_stub=self.operator_class.is_stub,
-            task_module=self.operator_class.__module__,
-            task_type=self.operator_class.__name__,
-            operator_name=operator_name,
-            dag=dag,
-            task_group=task_group,
-            start_date=start_date,
-            end_date=end_date,
-            disallow_kwargs_override=strict,
-            # For classic operators, this points to expand_input because kwargs
-            # to BaseOperator.expand() contribute to operator arguments.
-            expand_input_attr="expand_input",
-            # TODO: Move these to task SDK's BaseOperator and remove getattr
-            start_trigger_args=start_trigger_args,
-            start_from_trigger=start_from_trigger,
-        )
-        return op
+    def iterate(self, **mapped_kwargs: OperatorExpandArgument) -> 
IterableOperator:

Review Comment:
   Fixed in c038175a61 and 28050239bb. `.iterate()` and `.iterate_kwargs()` now 
mark the partial as expanded, so the "was never mapped" warning is gone (a test 
checks `recwarn`), and a dict value expands to its `(key, value)` pairs like 
`.expand()`. Since f55fb22809 the sync and async iteration paths share one 
helper for that rule.
   
   ---
   Drafted-by: Claude Fable 5.1; 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]

Reply via email to