dabla commented on code in PR #62922:
URL: https://github.com/apache/airflow/pull/62922#discussion_r4142422297
##########
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)
Review Comment:
Reproduced with your numbers: nothing was killed, and the sync items ran
their full 3 s. 3d8aed67ba does both things you suggested: `_run_tasks` calls
`on_kill()` before the executor cancels, and operators are registered in
`IndexedTaskRunner.run()`/`arun()`, so in the thread or coroutine that runs
them. Testing it turned up a second problem: the register was a set, and the
sub-operators of one iterated task compare equal (`BaseOperator.__eq__` looks
at `task_id`, `dag_id`, …), so it never held more than one of them, which also
limited the SIGTERM path. It is keyed by identity now, and each operator is
killed once although the runner calls `on_kill()` again after the timeout. The
item the timeout strikes directly is killed as it unwinds. The new tests run a
real 300 ms `execution_timeout` through `_run_execute_callable` for sync and
async items; the sync ones now stop after about 0.4 s.
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]