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


##########
task-sdk/src/airflow/sdk/definitions/batchedoperator.py:
##########
@@ -0,0 +1,556 @@
+#
+# 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
+from abc import ABCMeta, abstractmethod
+from collections.abc import Callable, Mapping, Sequence
+from typing import TYPE_CHECKING, Any, Generic, TypeVar
+
+import attrs
+
+from airflow.sdk import TriggerRule, timezone
+from airflow.sdk.bases.decorator import (
+    DecoratedMappedOperator,
+    FParams,
+    FReturn,
+    OperatorSubclass,
+    _TaskDecorator,
+    get_unique_task_id,
+)
+from airflow.sdk.bases.operator import (
+    BaseOperator,
+    coerce_resources,
+    coerce_timedelta,
+    get_merged_defaults,
+    parse_retries,
+)
+from airflow.sdk.bases.xcom import BaseXCom
+from airflow.sdk.definitions._internal.contextmanager import (
+    DagContext,
+    TaskGroupContext,
+)
+from airflow.sdk.definitions._internal.expandinput import (
+    EXPAND_INPUT_EMPTY,
+    DecoratedExpandInput,
+    DictOfListsExpandInput,
+    ExpandInput,
+    ListOfDictsExpandInput,
+    OperatorExpandArgument,
+    OperatorExpandKwargsArgument,
+)
+from airflow.sdk.definitions._internal.types import NOTSET
+from airflow.sdk.definitions.mappedoperator import (
+    MappedOperator,
+    OperatorPartial,
+    ensure_xcomarg_return_value,
+    prevent_duplicates,
+    validate_mapping_kwargs,
+)
+from airflow.sdk.definitions.xcom_arg import PlainXComArg, XComArg
+
+if TYPE_CHECKING:
+    from airflow.sdk.definitions.iterableoperator import IterableOperator, 
MappedIterableOperator
+    from airflow.sdk.definitions.mappedoperator import ValidationSource
+    from airflow.sdk.definitions.param import ParamsDict
+
+T = TypeVar("T", bound=OperatorPartial | _TaskDecorator)
+
+
+def validate_batch_size(size: int | XComArg) -> int | XComArg:
+    """
+    Validate the ``size`` handed to ``.batch()`` at DAG-definition time.
+
+    A literal size must be at least 2 (``.iterate()`` covers a single task 
instance). A runtime
+    size must be the return value of a plain, non-mapped task: the scheduler 
learns it from the
+    ``task_map`` row that the return value's push leaves behind (never from 
the XCom itself), so
+    a ``.map()``/``.filter()`` result, a pushed key or a mapped upstream 
cannot provide one.
+    """
+    if isinstance(size, PlainXComArg):
+        if size.operator.is_mapped:
+            raise ValueError(f"batch size cannot come from mapped task 
{size.operator.task_id!r}")
+        if size.key != BaseXCom.XCOM_RETURN_KEY:
+            raise ValueError(
+                f"batch size must be the return value of 
{size.operator.task_id!r}, not its {size.key!r} XCom"
+            )
+        return size
+    if isinstance(size, XComArg):
+        raise TypeError(f"batch size must be a plain XComArg, not 
{type(size).__name__}")
+    if size < 2:
+        raise ValueError(f"batch size must be at least 2, got {size}")
+    return size
+
+
[email protected](kw_only=True, repr=False)
+class BatchableOperator(Generic[T], metaclass=ABCMeta):
+    """
+    Intermediate abstraction for batched mapping.
+
+    This class decorates an OperatorPartial and stores configuration for 
batched mapping.
+    It is used to facilitate batched expansion of operators, allowing tasks to 
be mapped over batches
+    of data and then iterate over the batched data.
+
+    :param operator_partial: The partial operator to be batched.
+    :param size: The number of task instances to create. The input is 
distributed across them
+        round-robin (item ``i`` goes to task instance ``i % size``), not split 
into ``size``
+        contiguous chunks — this is *not* the same semantics as 
``itertools.batched(iterable, size)``.
+        See 
:class:`~airflow.sdk.definitions._internal.expandinput.BatchedExpandInput` for 
why
+        round-robin is used instead of contiguous chunking. Exactly ``size`` 
task instances are
+        always created; if the input yields fewer than ``size`` items, the 
surplus instances run
+        with no items and succeed immediately. May be an ``XComArg`` whose 
integer value is only
+        known at run time: the scheduler then creates that many task instances 
and each of them
+        resolves the same XCom to pick its share.
+    """
+
+    operator_partial: T
+    size: int | XComArg
+
+    @property
+    def operator_class(self) -> type[BaseOperator]:
+        return self.operator_partial.operator_class
+
+    @property
+    def kwargs(self) -> dict[str, Any]:
+        return self.operator_partial.kwargs
+
+    @abstractmethod
+    def iterate(self, **mapped_kwargs: OperatorExpandArgument) -> Any:
+        """
+        Iterate the operator over the provided mapped keyword arguments.
+
+        :param mapped_kwargs: Keyword arguments to expand against.
+        :return: An expanded operator or XComArg, depending on the subclass 
implementation.
+        """
+
+    @abstractmethod
+    def iterate_kwargs(self, kwargs: OperatorExpandKwargsArgument, *, strict: 
bool = True) -> Any:
+        """
+        Iterate the operator over a list of dictionaries or XComArg.
+
+        :param kwargs: List of dicts or XComArg to expand against.
+        :param strict: Whether to enforce strict argument checking.
+        :return: An expanded operator or XComArg, depending on the subclass 
implementation.
+        """
+
+    @abstractmethod
+    def _iterate(
+        self,
+        expand_input: ExpandInput,
+        *,
+        strict: bool,
+    ) -> IterableOperator | MappedIterableOperator:
+        """
+        Create an iterable operator for the given expansion input.
+
+        This method calls the _expand method first to get a MappedOperator 
based on expansion input,
+        then wraps it in either an IterableOperator or MappedIterableOperator 
depending on the batch size.
+
+        :param expand_input: The input to iterate against.
+        :param strict: Whether to enforce strict argument checking.
+        :return: An IterableOperator or MappedIterableOperator.
+        """
+
+    @abstractmethod
+    def _expand(
+        self,
+        expand_input: ExpandInput,
+        *,
+        strict: bool,
+        register_with_dag: bool = True,
+    ) -> MappedOperator:
+        """
+        Create a mapped operator for the given expansion input.
+
+        :param expand_input: The input to expand against.
+        :param strict: Whether to enforce strict argument checking.
+        :param register_with_dag: Whether to apply upstream relationships.
+        :return: A MappedOperator instance.
+        """
+
+
[email protected](kw_only=True, repr=False)
+class BatchedOperator(BatchableOperator[OperatorPartial]):
+    """
+    Concrete implementation of BatchableOperator for classic (non-decorated) 
operators.
+
+    This class wraps an OperatorPartial and provides batched expansion and 
iteration logic
+    for classic Airflow operators. It enables mapping tasks over batches of 
data, supporting
+    both direct expansion via keyword arguments and expansion via a list of 
dictionaries or XComArg.
+
+    :param operator_partial: The OperatorPartial instance to be batched and 
expanded.
+    :param size: The number of task instances to create for mapping. Items are 
distributed across
+        them round-robin (item ``i`` goes to task instance ``i % size``), not 
split into ``size``
+        contiguous chunks. Exactly ``size`` task instances are always created, 
even when the input
+        yields fewer items.
+    """
+
+    @property
+    def params(self) -> ParamsDict | dict:
+        return self.operator_partial.params
+
+    @property
+    def _expand_called(self) -> bool:
+        return self.operator_partial._expand_called
+
+    @_expand_called.setter
+    def _expand_called(self, value: bool) -> None:
+        self.operator_partial._expand_called = value
+
+    def iterate(self, **mapped_kwargs: OperatorExpandArgument) -> 
IterableOperator | MappedIterableOperator:
+        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.
+        expand_input = DictOfListsExpandInput(mapped_kwargs)
+        return self._iterate(expand_input, strict=False)
+
+    def iterate_kwargs(
+        self, kwargs: OperatorExpandKwargsArgument, *, strict: bool = True
+    ) -> IterableOperator | MappedIterableOperator:
+        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__}")
+
+        expand_input = ListOfDictsExpandInput(kwargs)
+        return self._iterate(expand_input, strict=strict)
+
+    def _iterate(
+        self,
+        expand_input: ExpandInput,
+        *,
+        strict: bool,
+    ) -> IterableOperator | MappedIterableOperator:
+        from airflow.sdk.definitions.iterableoperator import IterableOperator, 
MappedIterableOperator
+
+        # Unlike .expand(), neither 
OperatorPartial.iterate()/.iterate_kwargs() nor this class's own
+        # iterate()/iterate_kwargs() set _expand_called, so 
OperatorPartial.__del__ would otherwise
+        # warn "Task ... was never mapped!" even though 
.iterate()/.batch().iterate() legitimately
+        # consumed the partial.
+        self._expand_called = True
+        operator = self._expand(expand_input, strict=strict, 
register_with_dag=False)
+
+        if isinstance(self.size, XComArg) or self.size > 1:
+            return MappedIterableOperator(
+                mapped_operator=operator,
+                expand_input=expand_input,
+                batch_size=self.size,
+            )
+        return IterableOperator(
+            operator=operator,
+            expand_input=expand_input,
+        )
+
+    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
+
+        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)
+        )
+
+        try:
+            operator_name = self.operator_class.custom_operator_name  # type: 
ignore
+        except AttributeError:
+            operator_name = self.operator_class.__name__
+
+        return 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,
+            register_with_dag=register_with_dag,
+        )
+
+
[email protected](kw_only=True, repr=False)
+class DecoratedBatchedOperator(BatchableOperator[_TaskDecorator]):
+    """
+    Concrete implementation of BatchableOperator for decorated (TaskFlow) 
operators.
+
+    This class wraps a _TaskDecorator and provides batched expansion and 
iteration logic
+    for TaskFlow-style decorated Airflow operators. It enables mapping 
decorated tasks over
+    batches of data, returning XComArg objects for downstream dependencies and 
supporting
+    both direct expansion via keyword arguments and expansion via a list of 
dictionaries or XComArg.
+
+    :param operator_partial: The _TaskDecorator instance to be batched and 
expanded.
+    :param size: The number of task instances to create for mapping. Items are 
distributed across
+        them round-robin (item ``i`` goes to task instance ``i % size``), not 
split into ``size``
+        contiguous chunks. Exactly ``size`` task instances are always created, 
even when the input
+        yields fewer items.
+    """
+
+    @property
+    def is_setup(self) -> bool:
+        return self.operator_partial.is_setup
+
+    @property
+    def is_teardown(self) -> bool:
+        return self.operator_partial.is_teardown
+
+    @property
+    def function(self) -> Callable[FParams, FReturn]:
+        return self.operator_partial.function
+
+    @property
+    def operator_class(self) -> type[OperatorSubclass]:
+        return self.operator_partial.operator_class
+
+    @property
+    def multiple_outputs(self) -> bool:
+        return self.operator_partial.multiple_outputs
+
+    @property
+    def on_failure_fail_dagrun(self) -> bool:
+        return self.operator_partial.on_failure_fail_dagrun
+
+    def _validate_arg_names(self, func: ValidationSource, kwargs: dict[str, 
Any]):
+        self.operator_partial._validate_arg_names(func, kwargs)
+
+    @property
+    def returns_dag_result(self) -> bool:
+        return self.operator_partial.returns_dag_result
+
+    def iterate(self, **map_kwargs: OperatorExpandArgument) -> XComArg:
+        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 expand against")
+        self._validate_arg_names("expand", map_kwargs)

Review Comment:
   Fixed in f48db70ba0. The decorated path now passes `"iterate"` like the 
classic one, so the errors name the right method. It also stops applying 
`.expand()`'s mappable-type rule there, which had rejected 
`show.iterate(number=5)` with an "expand() got an unexpected type" error while 
the classic path accepts a scalar as a one-item input. A test pins both 
messages and the scalar case.
   
   ---
   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