amoghrajesh commented on code in PR #62922:
URL: https://github.com/apache/airflow/pull/62922#discussion_r4091719948
##########
airflow-core/tests/unit/models/test_mappedoperator.py:
##########
@@ -1828,3 +1828,88 @@ def test_mapped_operator_retry_delay_explicit(dag_maker):
# Should return the explicitly set value
assert mapped_deser.retry_delay == custom_retry_delay
+
+
[email protected](
+ ("batch_size", "items"),
+ [
+ pytest.param(5, [1, 2, 3], id="batch_size-larger-than-items"),
+ pytest.param(2, [1, 2, 3], id="batch_size-smaller-than-items"),
+ pytest.param(3, [1, 2, 3], id="batch_size-equal-to-items"),
+ pytest.param(2, 5, id="scalar-input-is-never-measured"),
+ ],
+)
+def test_batched_ti_count_is_batch_size_regardless_of_items(dag_maker,
session, batch_size, items):
+ from airflow.serialization.definitions.mappedoperator import
get_mapped_ti_count
+
+ with dag_maker(dag_id=f"test_batch_size_{batch_size}", session=session,
serialized=True) as dag:
+
MockOperator.partial(task_id="task").batch(size=batch_size).iterate(arg1=items)
Review Comment:
I guess its waiting for email to conclude but I assume the accepted name
(maybe `spread`) will be applied once agreed on devlist.
##########
task-sdk/src/airflow/sdk/execution_time/executor.py:
##########
@@ -0,0 +1,399 @@
+#
+# 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 threading
+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 TYPE_CHECKING, Any
+
+from airflow.sdk import BaseOperator, TaskInstanceState, timezone
+from airflow.sdk.definitions._internal.logging_mixin import LoggingMixin
+from airflow.sdk.exceptions import TaskDeferred
+from airflow.sdk.execution_time.context import set_current_context
+from airflow.sdk.execution_time.task_runner import (
+ _execute_async_task,
+ _execute_task,
+ _run_task_state_change_callbacks,
+)
+
+if TYPE_CHECKING:
+ from airflow.sdk import Context
+ from airflow.sdk.execution_time.task_runner import IndexedTaskInstance
+
+
+_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.
+ 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 map( # type: ignore[override]
Review Comment:
`map()` still yields in completion order while overriding `Executor.map`,
which promises submission order. The `# type: ignore[override]` now
acknowledges the mismatch instead of fixing it. Rename it, or stop subclassing
Executor.
##########
task-sdk/src/airflow/sdk/bases/xcom.py:
##########
@@ -564,3 +565,294 @@ def delete(
map_index=map_index,
),
)
+
+
+class XComIterable(Sequence):
+ """
+ An iterable that lazily fetches XCom values one by one instead of loading
all at once.
+
+ The class has two sides. The *producing* task builds it and grows it with
:meth:`append` /
+ :meth:`aappend`, each call pushing one more ``return_value_<index>`` XCom,
before returning it as
+ the task's result. Everything *downstream* (``.iterate()``, ``.expand()``,
a plain ``xcom_pull``)
+ only ever reads it, which is why the class implements the read-only
+ :class:`collections.abc.Sequence` rather than ``MutableSequence``: once
handed over it is a fixed
+ view of the values already pushed, and the two append methods are not part
of that contract.
+
+ Negative indices are not supported: every element is a remote fetch, and
resolving a negative
+ index against a lazily counted stream would cost a full walk just to find
the end.
+ """
+
+ def __init__(
+ self,
+ task_id: str,
+ dag_id: str,
+ run_id: str,
+ map_index: int | None = None,
+ length: int | None = None,
+ ):
+ 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._index = self.length
+
+ def __iter__(self) -> Iterator[Any]:
+ return _XComIterator(self)
+
+ def __len__(self) -> int:
+ return self.length
+
+ 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)]
+
+ if not (0 <= key < self.length):
Review Comment:
`XComIterable.__getitem__` still rejects negative indexes, so `result[-1]`
raises `IndexError`, while `FlattenedXComIterable.__getitem__` at line 761
accepts them. Same interface, different behaviour. Also still a Sequence
(read-only contract) with `append()` at 655 and `aappend()` at 675.
##########
task-sdk/src/airflow/sdk/definitions/_internal/expandinput.py:
##########
@@ -184,6 +363,85 @@ def iter_references(self) -> Iterable[tuple[Operator,
str]]:
if isinstance(x, XComArg):
yield from x.iter_references()
+ def iter_values(self, context: Mapping[str, Any]) -> Iterable[Any]:
+ from airflow.sdk.definitions.xcom_arg import XComArg
+
+ def _to_iterable(v: Any) -> Iterable:
Review Comment:
The sync and async `_to_iterable` already disagree. Sync returns `v.items()`
for a Mapping; async returns `list(v.items())`. Async also accepts `__aiter__`,
sync does not.
##########
airflow-core/src/airflow/serialization/definitions/mappedoperator.py:
##########
@@ -526,6 +587,25 @@ def _(task: SerializedBaseOperator | TaskSDKBaseOperator,
run_id: str, *, sessio
def _(task: SerializedMappedOperator | TaskSDKMappedOperator, run_id: str, *,
session: Session) -> int:
from airflow.serialization.serialized_objects import BaseSerialization,
_ExpandInputRef
+ def _get_parent_count() -> int:
+ if (group := task.get_closest_mapped_task_group()) is None:
+ return 1
+ return get_mapped_ti_count(group, run_id, session=session)
+
+ # See get_parse_time_mapped_ti_count: a batched task's count is fixed by
batch_size alone.
+ if isinstance(task, SerializedMappedOperator):
+ batch_size = task.resolve_batch_size(run_id, session=session)
+ elif isinstance(task.batch_size, int):
+ batch_size = task.batch_size
+ else:
+ # A runtime batch size lives in task_map and is only resolvable
through the serialized
+ # operator; SDK objects only reach here from tests that skip
serialization.
+ raise TypeError(
+ f"runtime batch size of {task.task_id!r} can only be resolved on a
serialized operator"
Review Comment:
`get_mapped_ti_count` raises `TypeError` when a SDK operator carries a
runtime batch size. The comment says "SDK objects only reach here from tests
that skip serialization." That assumes a dispatch function registered for
`TaskSDKMappedOperator`. If `dag.test()` or any non serialized path reaches it,
the user gets a confusing type error instead of a count. Worth confirming it is
genuinely unreachable or handling it.
##########
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:
still passes "expand" as the `ValidationSource`, so a bad argument to
`.iterate()` produces an error saying "expand". `ValidationSource` has an
"iterate" literal and `BatchedOperator.iterate` uses it.
##########
task-sdk/src/airflow/sdk/definitions/_internal/expandinput.py:
##########
@@ -41,6 +51,32 @@
OperatorExpandKwargsArgument = Union["XComArg", Sequence[Union["XComArg",
Mapping[str, Any]]]]
+async def aiterate(iterable: Any) -> AsyncIterator[Any]:
+ """
+ Iterate ``iterable`` from a coroutine without a blocking SDK call on the
loop thread.
+
+ An async iterable (``XComIterable``, ``LazyXComSequence``) is consumed
with ``async for``, so
+ its reads go through ``asend``. In-memory containers are iterated in
place. Anything else may
+ fetch on ``next()`` through a synchronous supervisor call, so each
``next()`` runs in a worker
+ thread: from there a blocking send waits for in-flight ``asend`` calls
instead of deadlocking
+ with them (see ``AsyncAwareExecutor.map``).
+ """
+ if hasattr(iterable, "__aiter__"):
+ async for item in iterable:
+ yield item
+ return
+
+ if isinstance(iterable, (list, tuple, set, frozenset, range, deque)):
+ for item in iterable:
+ yield item
+ return
+
+ iterator = iter(iterable)
+ sentinel = object()
+ while (item := await asyncio.to_thread(next, iterator, sentinel)) is not
sentinel:
Review Comment:
`aiterate` does one `asyncio.to_thread(next, ...)` per item for any sync
iterable that is not a known in-memory container. For a 1000 item iterate that
is 1000 thread dispatches. The reason is sound (a blocking send on the loop
thread would deadlock), but the cost should be in the docstring next to the
reason.
##########
task-sdk/src/airflow/sdk/definitions/iterableoperator.py:
##########
@@ -0,0 +1,453 @@
+#
+# 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 copy
+import os
+from collections import deque
+from collections.abc import Iterable, Mapping, Sequence
+from concurrent.futures import Future
+from itertools import chain
+from typing import TYPE_CHECKING, Any
+
+try:
+ # Python 3.12+
+ from itertools import batched # type: ignore[attr-defined]
+except ImportError:
+ from more_itertools import batched # type: ignore[no-redef]
+
+try:
+ # Python 3.11+
+ BaseExceptionGroup
+except NameError:
+ from exceptiongroup import BaseExceptionGroup
+
+from airflow.sdk import BaseXCom, TaskInstanceState, timezone
+from airflow.sdk.bases.operator import BaseOperator,
DecoratedDeferredAsyncOperator, event_loop
+from airflow.sdk.definitions._internal.expandinput import
PartitionedExpandInput
+from airflow.sdk.definitions.mappedoperator import MappedOperator
+from airflow.sdk.definitions.xcom_arg import MapXComArg, XComArg # noqa: F401
+from airflow.sdk.exceptions import (
+ AirflowRescheduleTaskInstanceException,
+ AirflowTaskTimeout,
+ TaskDeferred,
+)
+from airflow.sdk.execution_time.executor import ConcurrentExecutor,
TaskExecutor, collect_futures
+from airflow.sdk.execution_time.task_runner import IndexedTaskInstance
+
+if TYPE_CHECKING:
+ import jinja2
+
+ from airflow.sdk.definitions._internal.expandinput import ExpandInput
+ from airflow.sdk.definitions.context import Context
+ from airflow.sdk.execution_time.lazy_sequence import XComIterable
+
+
+class IterableOperator(BaseOperator):
+ """Object representing an iterable operator in a DAG."""
+
+ _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": 0, # We should not retry the IterableOperator,
only the mapped ti's should be retried
+ "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,
Review Comment:
`"execution_timeout": None` still strips the parent's wall clock limit.
Unchanged from last review. The docstring explains why sync sub tasks cannot be
timed out, but not that the parent's timeout is also discarded. A user setting
`execution_timeout=timedelta(minutes=5)` on a sync iterated task gets no
timeout anywhere.
##########
task-sdk/src/airflow/sdk/definitions/xcom_arg.py:
##########
@@ -180,6 +205,15 @@ 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:
Review Comment:
`aresolve` is added to the public `XComArg` base raising
`NotImplementedError`, with implementations on the built in subclasses. Anyone
who has subclassed `XComArg` gets an object that looks complete but blows up
the moment an iterated task touches it. Same shape for `iter_values` and
`aiter_values`. These need to be documented as new required surface, or given
working base implementations.
--
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]