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


##########
task-sdk/src/airflow/sdk/execution_time/task_runner.py:
##########
@@ -897,6 +923,337 @@ def mark_success_url(self) -> str:
         return self.log_url
 
 
+@dataclass
+class IndexedTaskState:
+    status: TaskInstanceState
+    result: Any | None = None
+    # Outlet asset events the sub-task recorded on a previous successful 
attempt. Outlet events
+    # only reach the server on the parent's success payload 
(_handle_current_task_success), so an
+    # attempt that fails never registers what its succeeded sub-tasks emitted. 
A sub-task skipped
+    # on retry (because it already succeeded) never re-executes, so it never 
re-emits into the
+    # fresh OutletEventAccessors created for the new attempt; persisting a 
snapshot here lets
+    # IterableOperator._run_task replay it, which is the only way those events 
survive, and it
+    # cannot double-emit because the failed attempt sent nothing.
+    outlet_events: list[dict[str, Any]] | None = None
+    # Digest of the input the sub-task ran with. An index only means the same 
work while its input
+    # is the same, so a checkpoint is honoured on a later attempt only when 
this still matches:
+    # after the upstream was cleared and produced other items, the sub-task 
runs again.
+    fingerprint: str | None = None
+    # The attempt that wrote the checkpoint. After a manual clear only the 
checkpoints written since
+    # are resumed from (see Checkpoints), which this tells apart from those 
left by the run before.
+    try_number: int = 0
+
+    @staticmethod
+    def build_key(index: int) -> str:
+        # The task state store is already scoped to the parent task instance 
(dag, run, task and
+        # map index), so the key carries no identity, only a namespace that 
keeps the operator's own
+        # entries apart from anything user code stores from inside a sub-task.
+        return f"_iterable_{index}"
+
+    def serialize(self) -> dict[str, Any]:
+        # The checkpoint travels to the supervisor as a JsonValue, which 
rejects anything that is not
+        # plain JSON (tuples, datetimes, models, ...). Serde turns those into 
JSON-compatible
+        # structures and restores them on read, exactly as XCom does with the 
same result.
+        data: dict[str, Any] = {"status": self.status.value}
+        if self.result is not None:
+            data["result"] = serde_serialize(self.result)
+        if self.outlet_events:
+            data["outlet_events"] = self.outlet_events
+        if self.fingerprint:
+            data["fingerprint"] = self.fingerprint
+        if self.try_number:
+            data["try_number"] = self.try_number
+        return data
+
+    @classmethod
+    def deserialize(cls, raw: Any) -> IndexedTaskState | None:
+        if not isinstance(raw, Mapping):
+            return None
+        return cls(
+            status=TaskInstanceState(raw["status"]),
+            result=serde_deserialize(raw.get("result")),
+            outlet_events=raw.get("outlet_events"),
+            fingerprint=raw.get("fingerprint"),
+            try_number=raw.get("try_number", 0),
+        )
+
+
+class IndexedTaskInstance(RuntimeTaskInstance):
+    """
+    Indexed task instance to run a mapped operator.
+
+    It shares the parent task instance's identity, so what an iteration pushes 
or stores lands in
+    the parent's scope, suffixed with the index so that iterations never 
overwrite each other:
+    XComs through :meth:`xcom_push`, task state through 
:attr:`task_state_store`. The operator's
+    own checkpoints go to the parent's store unsuffixed, under their 
``_iterable_<index>`` keys.
+    """
+
+    index: int
+    parent_task_state_store: TaskStateStoreAccessor
+    input_fingerprint: str | None = None
+
+    @classmethod
+    def create_indexed_task(
+        cls, *, context: Context, index: int, operator: BaseOperator, 
input_fingerprint: str | None = None
+    ) -> IndexedTaskInstance:
+        """
+        Create the runtime instance for one index of an iterated task from the 
parent's context.
+
+        The instance shares the parent task instance's identity (id, run, map 
index, try number and
+        retry budget), so XComs and task state land in the parent's scope, and 
carries the unmapped
+        operator for that index. The budget is the parent's ``max_tries`` 
rather than the operator's
+        ``retries``: a manual clear raises it, and it is the parent Airflow 
retries. The parent's
+        state store comes from the context, the same accessor the operator's
+        checkpoints use. ``model_construct`` skips Pydantic validation on 
purpose: one instance is built per
+        item, and the parent was validated already, so only the index needs 
checking here.
+        """
+        if index < 0:
+            raise ValueError(f"IndexedTaskInstance requires index >= 0, got 
{index}")
+        parent = context["ti"]
+        return cls.model_construct(
+            id=parent.id,
+            parent_task_state_store=context["task_state_store"],
+            task_id=operator.task_id,
+            dag_id=operator.dag_id,
+            run_id=parent.run_id,
+            map_index=parent.map_index,
+            index=index,
+            input_fingerprint=input_fingerprint,
+            max_tries=parent.max_tries,
+            start_date=parent.start_date,
+            state=TaskInstanceState.SCHEDULED.value,
+            is_mapped=True,
+            task=operator,
+            try_number=parent.try_number,
+        )
+
+    def xcom_push(
+        self,
+        key: str,
+        value: Any,
+    ):
+        super().xcom_push(key=f"{key}_{self.index}", value=value)
+
+    async def axcom_push(
+        self,
+        key: str,
+        value: Any,
+    ):
+        await super().axcom_push(key=f"{key}_{self.index}", value=value)
+
+    @cached_property
+    def task_state_store(self) -> TaskStateStoreAccessor:  # type: 
ignore[override]
+        """The parent's store seen from this iteration: keys are suffixed with 
the index."""
+        return IndexedTaskStateStoreAccessor(self.parent_task_state_store, 
self.index)
+
+    async def aget_state(self) -> IndexedTaskState | None:
+        return IndexedTaskState.deserialize(await 
self.parent_task_state_store.aget(self.state_key))
+
+    async def aset_state(self, state: IndexedTaskState) -> None:
+        await self.parent_task_state_store.aset(self.state_key, 
state.serialize())
+
+    @property
+    def is_async(self) -> bool:
+        return self.task.is_async
+
+    @property
+    def is_eligible_to_retry(self) -> bool:
+        """
+        Whether Airflow runs the parent again after this attempt fails.
+
+        The same rule the API server applies to the parent 
(``_is_eligible_to_retry``), so the
+        callbacks of an iteration agree with what happens to the task instance.
+        """
+        return self.max_tries != 0 and self.try_number <= self.max_tries
+
+    @property
+    def state_key(self) -> str:
+        return IndexedTaskState.build_key(self.index)
+
+    @property
+    def do_xcom_push(self) -> bool:
+        return self.task.do_xcom_push
+
+
+class IndexedTaskRunner(LoggingMixin):
+    """
+    Run one indexed task of an iterated task: its operator, against its own 
view of the context.
+
+    Named apart from Airflow's executors, which schedule task instances, and 
from the task runner
+    process, which runs the parent: this runs one index inside that process, 
sync or async.
+    """
+
+    def __init__(
+        self,
+        task_instance: IndexedTaskInstance,
+        active_operators: set[BaseOperator] | None = None,
+        active_operators_lock: threading.Lock | None = None,
+        outlet_events: OutletEventAccessors | None = None,
+    ):
+        """
+        Run an operator or trigger for one sub-task instance.
+
+        :param outlet_events: The accessor the sub-task's asset events are 
collected in, its own
+            so they can be checkpointed and merged apart from its siblings'. 
Created here when not
+            given; the caller reads it back through :attr:`outlet_events` 
after the run.
+        :param active_operators: Optional shared set the caller wants this 
operator registered
+            into for the duration of its execution (e.g. so 
IterableOperator.on_kill() can
+            propagate to whichever sub-tasks are currently in flight). 
Registration happens in
+            __enter__/__exit__ so callers no longer need their own try/finally 
bookkeeping.
+        :param active_operators_lock: Lock guarding ``active_operators``, 
required whenever
+            ``active_operators`` is given since multiple sub-tasks may run 
concurrently.
+        """
+        super().__init__()
+        self.task_instance = task_instance
+        self.outlet_events = outlet_events if outlet_events is not None else 
OutletEventAccessors()
+        self._result: Any | None = None
+        self._start_time: float | None = None
+        self._context: Context | None = None
+        self._active_operators = active_operators
+        self._active_operators_lock = active_operators_lock
+
+    @property
+    def dag_id(self) -> str:
+        return self.task_instance.dag_id
+
+    @property
+    def task_id(self) -> str:
+        return self.task_instance.task_id
+
+    @property
+    def task_index(self) -> int:
+        return self.task_instance.index
+
+    @property
+    def operator(self) -> BaseOperator:
+        return self.task_instance.task
+
+    @property
+    def is_async(self) -> bool:
+        return self.task_instance.is_async
+
+    @contextmanager
+    def indexed_context(self, context: Context) -> Iterator[Context]:
+        """
+        Enter the parent's context as this indexed task sees it.
+
+        Yields a copy of the parent's context with this task's own task 
instance, its indexed view
+        of the task state store and its own outlet events, remembered on the 
runner and made the
+        current context for the duration of the block. The parent's context is 
left untouched:
+        ``context_update_for_unmapped`` sets ``ti.task`` on whatever ``ti`` it 
finds, which must be
+        this task's, not the parent's.
+        """
+        indexed_context: Context = {
+            **clone_context(context),
+            "ti": self.task_instance,
+            "task_instance": self.task_instance,
+            "task_state_store": self.task_instance.task_state_store,
+            "outlet_events": self.outlet_events,
+        }
+        self._context = indexed_context
+        with set_indexed_context(indexed_context):
+            yield indexed_context
+
+    def run(self, context: Context):
+        """Run the operator synchronously against this indexed task's own view 
of ``context``."""
+        with self.indexed_context(context) as indexed_context:
+            return _execute_task(indexed_context, self.task_instance, self.log)
+
+    async def arun(self, context: Context):
+        """Run the async operator against this indexed task's own view of 
``context``."""
+        with self.indexed_context(context) as indexed_context:
+            return await _execute_async_task(indexed_context, 
self.task_instance, self.log)
+
+    def __enter__(self):
+        self._start_time = time.monotonic()
+
+        if self._active_operators is not None and self._active_operators_lock 
is not None:
+            with self._active_operators_lock:
+                self._active_operators.add(self.operator)
+
+        if self.log.isEnabledFor(logging.INFO):
+            self.log.info(
+                "Running attempt %s of %s for %s with index %s in %s mode.",
+                self.task_instance.try_number,
+                self.task_instance.max_tries + 1,
+                self.task_instance.task_id,
+                self.task_index,
+                "async" if self.is_async else "sync",
+            )
+        return self
+
+    def __exit__(self, exc_type, exc_value, traceback):
+        if self._active_operators is not None and self._active_operators_lock 
is not None:
+            with self._active_operators_lock:
+                self._active_operators.discard(self.operator)
+
+        elapsed = time.monotonic() - self._start_time if self._start_time else 
0.0
+
+        if exc_value:
+            # Cancelled because the task is stopping, for a reason another 
iteration raised: this
+            # iteration neither failed nor will be retried on its own account, 
so it gets no state
+            # and no callback. The iteration that stopped the task reports its 
own outcome.
+            if isinstance(exc_value, CancelledError):
+                raise exc_value
+            # Non-Exception BaseExceptions (e.g. DeadlockImminentError,
+            # KeyboardInterrupt, SystemExit) must never be retried: they
+            # signal conditions where continuing is meaningless.
+            # Re-raise immediately without retry. The parent's 
execution_timeout is the exception:
+            # it strikes this iteration because it runs on the main thread, 
and the task is retried
+            # for it like for any other error, so it takes the retry decision 
below.
+            if not isinstance(exc_value, (Exception, AirflowTaskTimeout)):
+                self.task_instance.state = TaskInstanceState.FAILED
+                if self._context is not None:
+                    _run_task_state_change_callbacks(
+                        self.task_instance.task, "on_failure_callback", 
self._context, self.log
+                    )
+                raise exc_value
+            if isinstance(exc_value, AirflowSkipException):
+                self.task_instance.state = TaskInstanceState.SKIPPED
+                if self._context is not None:
+                    _run_task_state_change_callbacks(
+                        self.task_instance.task, "on_skipped_callback", 
self._context, self.log
+                    )
+                raise exc_value
+            # AirflowFailException fails the parent without a retry, whatever 
budget is left.
+            if isinstance(exc_value, AirflowFailException) or not 
self.task_instance.is_eligible_to_retry:
+                self.log.error(
+                    "Task instance %s for %s failed on attempt %s in %.2f 
seconds due to: %s",
+                    self.task_index,
+                    self.task_instance.task_id,
+                    self.task_instance.try_number,
+                    elapsed,
+                    exc_value,
+                )
+                self.task_instance.state = TaskInstanceState.FAILED
+                if self._context is not None:
+                    _run_task_state_change_callbacks(
+                        self.task_instance.task, "on_failure_callback", 
self._context, self.log
+                    )
+                raise exc_value
+            self.task_instance.end_date = datetime.now(tz=timezone.utc)
+            self.task_instance.state = TaskInstanceState.UP_FOR_RETRY

Review Comment:
   Reproduced your example: success, retry, failure, and the task FAILED. 
7758e0e7a1 postpones a failed item's failure/retry callback until every item 
has run: `__exit__` now only notes the failure, and `_run_tasks` decides the 
task's fate from the exception it hands the runner (fail-fast, a policy FAIL, 
or no attempts left means no retry) and then reports each failed item with the 
matching callback, one after another. The exceptions `_run_tasks` rejects 
therefore count as final. Success and skip callbacks still fire right away. 
Failures no item owns (input resolution, items cancelled by the parent's 
timeout) fire none, since the iterated task carries no callbacks of its own; 
the page now has a "Callbacks" section saying which callbacks run when.
   
   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]

Reply via email to