kaxil commented on code in PR #74349:
URL: https://github.com/apache/airflow/pull/74349#discussion_r4198236594
##########
airflow-core/src/airflow/models/taskinstance.py:
##########
@@ -1212,14 +1260,45 @@ def is_premature(self) -> bool:
# is the task still in the retry waiting period?
return self.state == TaskInstanceState.UP_FOR_RETRY and not
self.ready_for_retry()
+ def is_replaced_generated_work(self, *, session: Session) -> bool:
+ """Whether a pre-region mapped expansion was cleared as a whole and
now has a replacement region."""
+ if self.region_id != SENTINEL_REGION_ID or self.region_index < 0:
+ return False
+ return (
+ session.scalar(
+ select(DynamicRegion.id)
+ .where(
+ DynamicRegion.dag_id == self.dag_id,
+ DynamicRegion.run_id == self.run_id,
+ DynamicRegion.node_id == self.task_id,
+ DynamicRegion.parent_region_id.is_(None),
+ DynamicRegion.forked_from_region_id.is_(None),
+ )
+ .limit(1)
+ )
+ is not None
+ )
+
def archive(self, *, reason: str, session: Session) -> None:
"""Remove this attempt from the working set while retaining its UUID
and children."""
- current = session.scalar(
- select(TaskInstance.working_set).where(TaskInstance.id ==
self.id).with_for_update()
- )
- if current is not True:
+ current = session.execute(
+ select(
+ TaskInstance.working_set,
+ TaskInstance.state,
+ TaskInstance.start_date,
+ TaskInstance.end_date,
+ )
+ .where(TaskInstance.id == self.id)
+ .with_for_update()
+ ).one_or_none()
+ if current is None or current.working_set is not True:
raise ValueError("An archived task instance cannot be archived
again")
- if self.state not in State.finished:
+ state: Any = inspect(self)
+ for name in ("state", "start_date", "end_date"):
+ if not state.attrs[name].history.has_changes():
+ setattr(self, name, getattr(current, name))
+ never_started = reason == "superseded" and self.start_date is None
Review Comment:
`start_date is None` doesn't hold for most rows that never ran.
`prepare_db_for_next_try` copies every column into the successor, so a cleared
successor (state None) or a pending `UP_FOR_RETRY` try still carries the
previous try's `start_date` and `end_date`. If one of those gets superseded
(say `consume` in iteration 3 was cleared on its own, and then the gate of
iteration 2 is cleared with later iterations), it lands here as FAILED with the
old try's timestamps, which reads as a real failed run in the try history. The
clear path above already treats `None` and `UP_FOR_RETRY` as "hasn't run" (line
491). Could this decide from the refreshed state instead, something like
`self.state in (None, SCHEDULED, QUEUED, UP_FOR_RETRY)`? A test that supersedes
a cleared successor would pin it.
##########
airflow-core/src/airflow/ti_deps/deps/loop_archival_dep.py:
##########
@@ -0,0 +1,32 @@
+# 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
+
+from airflow.ti_deps.deps.base_ti_dep import BaseTIDep
+
+
+class LoopArchivalDep(BaseTIDep):
+ """A rerun gate waits for the generated work it replaces to terminate."""
+
+ NAME = "Loop archival"
+ IS_TASK_DEP = True
+
+ def _get_dep_statuses(self, ti, dep_context, session):
+ from airflow.models.loop_clear import loop_gate_waits_for_archival
+
+ if loop_gate_waits_for_archival(ti, session=session):
+ yield self._failing_status(reason="Superseded loop executions are
still terminating")
Review Comment:
`DagRun.update_state` can't tell this wait apart from a dependency that will
never be met, and `_are_premature_tis` doesn't neutralise it the way
`ignore_in_retry_period` does for retries. That bites when the RESTARTING
later-pass row has an upstream inside its pass. Take a body of `prepare >>
process` with `process[3]` running: clearing the gate of iteration 1 archives
`prepare[3]` and moves `process[3]` to RESTARTING. On the next scheduling pass
the rerun gate fails only on this dep, `process[3]` fails its trigger rule
because its upstream is archived, and the task after the loop waits on the
gate, so the run goes FAILED with `all_tasks_deadlocked` before the worker acks
the termination. The rewind test misses it because its loop body is a single
task with no upstream. Could this wait count as premature (a `DepContext` flag
that `_are_premature_tis` sets), or could `update_state` treat an unfinished
RESTARTING TI as keeping the run alive?
##########
airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py:
##########
@@ -197,33 +221,107 @@ def _patch_ti_validate_request(
session: SessionDep,
map_index: int | None = -1,
update_mask: list[str] | None = None,
+ *,
+ lock: bool = True,
) -> tuple[SerializedDAG, list[TI], dict]:
+ _validate_region_selection(body)
dag = get_latest_version_of_dag(dag_bag, dag_id, session)
- if not dag.has_task(task_id):
+ if lock:
+ _lock_patch_runs(dag, dag_run_id, body, session)
+ if body.region_id is None and not dag.has_task(task_id):
raise HTTPException(status.HTTP_404_NOT_FOUND, f"Task '{task_id}' not
found in Dag '{dag_id}'")
query = (
select(TI)
.where(TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id ==
task_id)
.options(joinedload(TI.rendered_task_instance_fields))
)
- if map_index is not None:
- query = query.where(TI.map_index == map_index)
- else:
- query = query.order_by(TI.map_index)
-
+ if body.region_id is None:
+ resolver = TaskCoordinateResolver(dag_bag, session)
+ for version in session.scalars(
Review Comment:
`dag_version_id` can be NULL here for task instances from runs that predate
Dag versioning (migration 0047 added the column as nullable with no backfill,
and those runs have no `created_dag_version_id` either). With `version=None`,
`get_task` falls through to `_producer_task`, which raises a bare
`ValueError("Pinned Dag ... not found")`, and nothing on the PATCH or bulk path
catches it, so adding a note to or marking such a TI now returns a 500. A TI
with no version predates loops, so could this query just skip it with
`.where(TI.dag_version_id.is_not(None))`? `select_loop_clear_scope` makes the
same `get_task(..., dag_version_id=None)` call once the Dag has a loop, which
turns a clear of such a run into a 409.
##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -2754,10 +2870,18 @@ def _flush_ti_buffer(*, drain: bool = False) -> int:
while len(ti_carry) >= _TI_CHUNK_SIZE:
slice_tis = ti_carry[:_TI_CHUNK_SIZE]
del ti_carry[:_TI_CHUNK_SIZE]
- clear_task_instances(slice_tis, session=session)
+ clear_task_instances_for_runs(
Review Comment:
These slices cut across runs, and the buffer query has no ordering, so one
run's TIs can be split over two calls. For a loop Dag that breaks: if the first
slice holds a gate of run B, `clear_task_instances_for_runs` archives every
later pass of B, including rows still waiting in `ti_carry` for the next slice.
The next call then raises "An archived task instance cannot be cleared" or
"...no longer live", and the whole partition clear rolls back. Two runs of a
10-iteration loop with about 30 tasks per pass is enough to cross
`_TI_CHUNK_SIZE`. Could the flush keep each run's TIs in a single call, for
example by grouping `ti_carry` by run_id and only flushing whole runs?
##########
airflow-core/src/airflow/cli/commands/dag_command.py:
##########
@@ -234,7 +234,7 @@ def _bulk_clear_runs(
tis = session.scalars(ti_query).all()
if not tis:
continue
- clear_task_instances(list(tis), session=session)
+ clear_task_instances_for_runs(tis, session=session)
Review Comment:
This passes neither of the knobs that `SerializedDAG.clear` sets. So
`airflow dags clear --only-failed` still archives the later passes of a failed
gate (`dag.clear` passes `later_loop_iterations=not (only_failed or
only_running)`), and a clear with no state filter gives a pre-region mapped
expansion successor tries in place, where `dag.clear` and the partition clear
archive it and mint a region. Should this mirror `dag.clear`, e.g.
`later_loop_iterations=not state_filter` and `whole_task_keys={(ti.dag_id,
ti.run_id, ti.task_id) for ti in tis} if not state_filter else ()`?
##########
airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py:
##########
@@ -484,17 +733,35 @@ def _categorize_task_instances(
# Filter at database level using exact tuple matching instead of
fetching all combinations
# and filtering in Python
task_keys_list = list(task_keys)
- query = select(TI).where(tuple_(TI.dag_id, TI.run_id, TI.task_id,
TI.map_index).in_(task_keys_list))
-
- task_instances = self.session.scalars(query).all()
- task_instances_map = {
- (ti.dag_id, ti.run_id, ti.task_id, ti.map_index if ti.map_index is
not None else -1): ti
- for ti in task_instances
- }
+ public_index = public_map_index_expression(TI)
+ query = select(TI, public_index).where(
+ tuple_(TI.dag_id, TI.run_id, TI.task_id,
public_index).in_(task_keys_list)
+ )
+ rows = self.session.execute(query).all()
+ self._reject_unscoped_loop_tasks([ti for ti, _ in rows])
+ task_instances_map = {}
Review Comment:
The fourth element of this row value is now the `CASE ... EXISTS`
expression, and the row `IN` is the only predicate. MySQL only range-optimizes
a row-constructor `IN` when every element is a plain column, so there this
likely scans `task_instance` and runs the correlated `dynamic_region` lookup
per row, on every bulk PATCH or DELETE with an explicit map_index. Adding a
plain-column prefilter, e.g. `.where(tuple_(TI.dag_id, TI.run_id,
TI.task_id).in_({key[:3] for key in task_keys_list}))`, would let it use the
unique key and keep the exact match.
##########
airflow-core/src/airflow/models/loop_clear.py:
##########
@@ -0,0 +1,449 @@
+# 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
+
+from collections.abc import Collection, Iterator
+from dataclasses import dataclass
+from typing import TYPE_CHECKING, Literal
+from uuid import UUID
+
+from sqlalchemy import exists, select, tuple_
+
+from airflow.exceptions import AirflowClearRunningTaskException
+from airflow.models.dagbag import DBDagBag
+from airflow.models.dynamic_region import (
+ SENTINEL_REGION_ID,
+ DynamicRegion,
+ load_region_ancestry,
+ loop_position,
+)
+from airflow.models.task_coordinates import TaskCoordinateResolver,
enclosing_loop
+from airflow.models.taskinstance import TaskInstance,
_get_relevant_map_indexes, clear_task_instances
+from airflow.serialization.definitions.mappedoperator import
get_mapped_ti_count
+from airflow.utils.state import DagRunState, TaskInstanceState
+
+if TYPE_CHECKING:
+ from sqlalchemy.orm import Session
+
+ from airflow.serialization.definitions.taskgroup import
SerializedLoopTaskGroup
+
+
+@dataclass(frozen=True)
+class LoopClearScope:
+ """Live executions to retry and generated executions to archive."""
+
+ retry_ids: frozenset[UUID]
+ archive_ids: frozenset[UUID]
+
+
+def select_loop_clear_scope(
+ selected: Collection[TaskInstance],
+ *,
+ whole_expansion_ids: Collection[UUID] = (),
+ upstream: bool = False,
+ downstream: bool = True,
+ later_loop_iterations: bool = True,
+ session: Session,
+) -> LoopClearScope:
+ """Select loop clear executions; mutation callers must hold the DagRun
lock."""
+ if not selected:
+ return LoopClearScope(frozenset(), frozenset())
+ runs = {(ti.dag_id, ti.run_id) for ti in selected}
+ if len(runs) != 1:
+ raise ValueError("Loop clear selection must belong to one DagRun")
+ dag_id, run_id = runs.pop()
+ live = {
+ ti.id: ti
+ for ti in session.scalars(
+ select(TaskInstance)
+ .where(
+ TaskInstance.dag_id == dag_id,
+ TaskInstance.run_id == run_id,
+ TaskInstance.working_set.is_(True),
+ )
+ .execution_options(populate_existing=True)
+ )
+ }
+ selected_ids = {ti.id for ti in selected}
+ if not selected_ids <= live.keys():
+ raise ValueError("Loop clear selection contains executions that are no
longer live")
+ if not set(whole_expansion_ids) <= selected_ids:
+ raise ValueError("Whole-task selection requires an explicitly selected
execution")
+ resolver = TaskCoordinateResolver(DBDagBag(), session)
+ regions = load_region_ancestry(
+ {ti.region_id for ti in live.values()}, dag_id=dag_id, run_id=run_id,
session=session
+ )
+ retry = {ti_id: live[ti_id] for ti_id in selected_ids}
+ for ti in tuple(retry.values()):
+ task = resolver.get_task(dag_id, run_id, ti.task_id,
dag_version_id=ti.dag_version_id)
+ group = enclosing_loop(task)
+ if group is not None and loop_position(regions, ti.region_id,
ti.region_index, group.node_id) is None:
+ raise ValueError("Loop clear selection requires an execution
inside its pinned loop")
+ if ti.id in whole_expansion_ids:
+ if not task.get_needs_expansion():
+ raise ValueError("Whole-expansion selection requires a mapped
task")
+ retry.update(
+ (other.id, other)
+ for other in resolver.resolve(dag_id=dag_id, run_id=run_id,
task_id=ti.task_id, caller=ti)
+ )
+ for is_upstream, enabled in ((True, upstream), (False, downstream)):
+ if not enabled:
+ continue
+ for ti in tuple(retry.values()):
+ task = resolver.get_task(dag_id, run_id, ti.task_id,
dag_version_id=ti.dag_version_id)
+ contexts = resolver.producer_contexts(ti)
+ count = (
+ get_mapped_ti_count(task, run_id, session=session,
producer_contexts=contexts)
+ if task.get_needs_expansion() and ti.region_index >= 0
+ else None
+ )
+ for relative in task.get_flat_relatives(upstream=is_upstream):
+ indexes = _get_relevant_map_indexes(
+ task=task,
+ run_id=run_id,
+ map_index=resolver.public_map_index(ti),
+ relative=relative,
+ ti_count=count,
+ session=session,
+ producer_contexts=contexts,
+ )
+ relative_loop = enclosing_loop(relative)
+ task_loop = enclosing_loop(task)
+ matches: Collection[TaskInstance]
+ if relative_loop is not None and (
+ task_loop is None or task_loop.group_id !=
relative_loop.group_id
+ ):
+ matches = [
+ other
+ for other in live.values()
+ if other.task_id == relative.task_id
+ and (
+ indexes is None
+ or (isinstance(indexes, int) and
resolver.public_map_index(other) == indexes)
+ or (isinstance(indexes, range) and
resolver.public_map_index(other) in indexes)
+ )
+ ]
+ else:
+ matches = resolver.resolve(
Review Comment:
Every selected TI walks its relatives one at a time here, and each `resolve`
runs at least two queries, with `get_mapped_ti_count` and `producer_contexts`
on top. Since the route sends any Dag that contains a loop down this path,
clearing an ordinary 1000-wide mapped task with downstream in such a Dag costs
roughly 1000 x (1 + 2 relatives x 2) queries that keep returning the same
downstream rows, all under the DagRun lock. With `only_failed`/`only_running`
the route also drops the SQL state filter and plans twice. Could relatives be
expanded once per `(task_id, region_id, dag_version_id)` when the task and
relative share no mapped group (`_get_relevant_map_indexes` returns None there
whatever the index), and `get_mapped_ti_count` be computed once per task and
region?
##########
airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_instances.py:
##########
@@ -896,7 +908,7 @@ def post_clear_task_instances(
error_message = f"Dag Run id {dag_run_id} not found in dag
{dag_id}"
raise HTTPException(status.HTTP_404_NOT_FOUND, error_message)
# Get the specific dag version:
- dag = get_dag_for_run(dag_bag, dag_run, session)
+ dag = get_dag_for_run_or_latest_version(dag_bag, dag_run, dag_id,
session)
Review Comment:
On an unversioned bundle (the default dags folder), the scheduler keeps
moving unfinished TIs to the latest version and `verify_integrity` adds tasks
that appear mid-run, so a running run's live structure can be newer than
`created_dag_version_id`. With the creation version, the non-loop path now
computes relatives and task groups from the old definition: if `a >> b` is
edited to `a >> b >> c` while the run is going, clearing `b` with downstream no
longer clears `c`, and a `task_group_id` added in the edit returns 404. The
loop path already resolves each TI against its own pinned version, so is the
creation version needed here, or could only the `loop_aware` probe look at it?
##########
airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py:
##########
@@ -1574,6 +1670,168 @@ def body():
ready =
dr.task_instance_scheduling_decisions(session=session).schedulable_tis
assert {(ti.task_id, ti.region_index) for ti in ready} ==
{("body.left", 1), ("body.right", 1)}
+ @pytest.fixture
+ def started_loop_successor(self, client, session, running_loop_gate):
+ gate = running_loop_gate
+ dr = gate.dag_run
+ response = client.patch(
+ f"/execution/task-instances/{gate.id}/state",
+ json={"state": "success", "end_date":
DEFAULT_END_DATE.isoformat()},
+ )
+ assert response.status_code == 204
+ session.expire_all()
+ running = next(
+ ti
+ for ti in dr.get_task_instances(session=session)
+ if ti.task_id != gate.task_id and ti.region_index == 1
+ )
+ running.state = State.QUEUED
+ for ti in dr.get_task_instances(session=session):
+ if ti.task_id != gate.task_id and ti.region_index == 0:
+ ti.state = State.SUCCESS
+ session.commit()
+ running_id = running.id
+ response = client.patch(
+ f"/execution/task-instances/{running_id}/run",
+ json={
+ "state": "running",
+ "hostname": "archiving-worker",
+ "unixname": "airflow",
+ "pid": 123,
+ "start_date": DEFAULT_START_DATE.isoformat(),
+ },
+ )
+ assert response.status_code == 200, response.text
+ session.expire_all()
+ return gate, running_id
+
+ def
test_loop_rewind_waits_for_worker_acknowledgement_before_gate_admission(
+ self, client, session, started_loop_successor, dag_maker
+ ):
+ gate, running_id = started_loop_successor
+ dr = gate.dag_run
+ running = session.get(TaskInstance, running_id)
+
+ for _ in range(2):
+ clear_loop_task_instances([session.get(TaskInstance, gate.id)],
downstream=False, session=session)
+ session.commit()
+ session.expire_all()
+ assert running.id == running_id
+ assert running.state == State.RESTARTING
+ assert running.working_set is True
+ dr.dag = dag_maker.serialized_dag
+ decision = dr.task_instance_scheduling_decisions(session=session)
+ assert gate.id not in {ti.id for ti in decision.schedulable_tis}
+ dr.update_state(session=session)
+ assert dr.state == DagRunState.RUNNING
+ session.commit()
+ gate = next(
+ ti
+ for ti in dr.get_task_instances(session=session)
+ if ti.task_id == gate.task_id and ti.region_index ==
gate.region_index
+ )
+
+ response = client.patch(
+ f"/execution/task-instances/{running_id}/state",
+ json={
+ "state": "server_terminated",
+ "end_date": DEFAULT_END_DATE.isoformat(),
+ "hostname": "archiving-worker",
+ "pid": 123,
+ },
+ )
+
+ assert response.status_code == 204, response.text
+ session.expire_all()
+ archived = session.get(TaskInstance, running_id)
+ assert (archived.working_set, archived.archived_reason) == (None,
"superseded")
+ dr.dag = dag_maker.serialized_dag
+ decision = dr.task_instance_scheduling_decisions(session=session)
+ assert gate.id in {ti.id for ti in decision.schedulable_tis}
+ assert not any(ti.region_index == 1 for ti in
dr.get_task_instances(session=session))
+
+ def test_restart_ack_locks_run_before_task(self, client, session,
create_task_instance, mocker):
+ if session.bind.dialect.name == "sqlite":
+ pytest.skip("SQLite has no row locks")
+ ti = create_task_instance(state=State.QUEUED)
Review Comment:
`create_task_instance` defaults to map_index -1 in the sentinel region, and
`ti_update_state` only takes the DagRun lock for regional rows (`in_region`).
So on Postgres and MySQL the ack's first `FOR UPDATE` is the TI lock, `assert
request_session.info.get("run_locked")` fails inside the ack thread, and
`ack.result()` re-raises it.
`test_stopped_report_locks_dag_run_only_for_regional_task` asserts the opposite
for `(-1, False)`. Making this TI regional (a loop member or a mapped index)
would exercise the path that actually takes the run lock.
##########
airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_instances.py:
##########
@@ -1144,14 +1315,28 @@ def patch_task_instance_dry_run(
session: SessionDep,
map_index: int | None = None,
update_mask: list[str] | None = Query(None),
+ region_id: UUID | None = None,
+ region_index: int | None = None,
) -> TaskInstanceCollectionResponse:
"""Update a task instance dry_run mode."""
tis: Sequence[TI]
+ body = patch_region_selection(body, region_id, region_index)
dag, tis, data = _patch_ti_validate_request(
- dag_id, dag_run_id, task_id, dag_bag, body, session, map_index,
update_mask
+ dag_id, dag_run_id, task_id, dag_bag, body, session, map_index,
update_mask, lock=False
Review Comment:
This dry run now goes through the same loop checks as the real PATCH, so it
can answer 409 ("Select a region and index for this loop task", or an ambiguous
producer), but its decorators still document only 404 and 400. Adding
`HTTP_409_CONFLICT` to these two dry-run decorators and the task group one
would keep the spec and generated clients in line with the non-dry-run siblings.
##########
airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py:
##########
@@ -545,13 +812,92 @@ def handle_bulk_create(
}
)
+ def _handle_regional_bulk(
+ self, action: MutationAction, results: BulkActionResponse
+ ) -> tuple[MutationAction, set[tuple[str, str, str, int]], set[tuple[str,
str, str]]]:
+ regional = [
+ entity
+ for entity in action.entities
+ if isinstance(entity, BulkTaskInstanceBody)
+ and (entity.region_id is not None or entity.region_index is not
None)
+ ]
+ deleting = isinstance(action, BulkDeleteAction)
+ specific, whole = self._categorize_entities(
+ action.entities, results, method="DELETE" if deleting else "PUT",
action_name=action.action.value
+ )
+ keys = {key[:3] for key in specific} | whole
+ run_keys = {key[:2] for key in keys}
+ for entity in action.entities:
+ if isinstance(entity, BulkTaskInstanceBody) and
(entity.include_future or entity.include_past):
+ dag_id, run_id, task_id, _ =
self._extract_task_identifiers(entity)
+ if (dag_id, run_id, task_id) in keys:
+ dag = get_latest_version_of_dag(self.dag_bag, dag_id,
self.session)
+ run_keys.update(
+ (dag_id, selected_run)
+ for selected_run in get_run_ids(
+ dag, run_id, entity.include_future,
entity.include_past, session=self.session
+ )
+ )
+ self.session.scalars(
+ select(DagRun)
+ .where(
+ tuple_(DagRun.dag_id, DagRun.run_id).in_(run_keys),
+ )
+ .order_by(DagRun.dag_id, DagRun.run_id)
+ .with_for_update()
+ ).all()
+ for entity in regional:
+ dag_id, run_id, task_id, map_index =
self._extract_task_identifiers(entity)
+ if (dag_id, run_id, task_id) not in keys:
+ continue
+ try:
+ dag, tis, data = _patch_ti_validate_request(
+ dag_id,
+ run_id,
+ task_id,
+ self.dag_bag,
+ entity,
+ self.session,
+ map_index,
+ getattr(action, "update_mask", None),
+ )
+ if deleting:
+ for ti in tis:
+ self.session.delete(ti)
Review Comment:
This removes only the live row. The non-regional delete uses
`TI.delete_attempts`, which also removes archived tries, so deleting a loop
pass by region leaves its earlier tries (and the logs and XCom keyed by their
attempt UUIDs) behind, and `get_last_try_numbers` keeps counting them if that
coordinate comes back. Could this delete by coordinate with
`include_all_attempts=True`, the way `delete_attempts` does?
##########
airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_instances.py:
##########
@@ -1010,14 +1156,18 @@ def _collect_relatives(run_id: str, direction:
Literal["upstream", "downstream"]
if body.note is not None:
_patch_task_instance_note(
task_instance_body=body,
- tis=task_instances,
+ tis=list(task_instances),
user=user,
)
+ # The reload below refreshes with populate_existing, which would
discard an unflushed note.
+ session.flush()
- task_instances = _reload_tis_with_rendered_fields(task_instances, session)
+ task_instances = _reload_tis_with_rendered_fields(list(task_instances),
session)
Review Comment:
The rows this clear archived come back from `clear_task_instances`, but this
reload is `select(TI).where(TI.id.in_(...))`, which isn't a primary-key lookup,
so `_restrict_to_current_attempts` adds `working_set IS TRUE` and every
archived row drops out of the response. After clearing a gate with later
iterations, the dry run lists those passes and the real response doesn't, and
`total_entries` differs.
`test_legacy_whole_clear_retains_affected_execution_and_note` (whole=True) and
the `real >= reported` check both expect them, so either those fail or
something I didn't find keeps archived rows in. Since the ids are already
explicit, passing `include_all_attempts=True` in
`_reload_tis_with_rendered_fields` would cover it.
##########
airflow-core/src/airflow/api_fastapi/core_api/services/public/task_coordinates.py:
##########
@@ -0,0 +1,113 @@
+# 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
+
+from dataclasses import dataclass
+from typing import TYPE_CHECKING, Any, TypeVar
+from uuid import UUID
+
+from fastapi import HTTPException, status
+from pydantic import BaseModel
+
+from airflow._shared.state import TaskScope
+from airflow.models.dynamic_region import SENTINEL_REGION_ID,
AmbiguousProducerError
+from airflow.models.task_coordinates import TaskCoordinateResolver
+
+if TYPE_CHECKING:
+ from sqlalchemy.orm import Session
+
+ from airflow.models.dagbag import DBDagBag
+ from airflow.models.task_coordinates import TaskCoordinate
+
+
+Response = TypeVar("Response", bound=BaseModel)
+
+
+@dataclass
+class TaskCoordinateView:
+ """Present public mapping coordinates without changing ORM state."""
+
+ value: TaskCoordinate
+ resolver: TaskCoordinateResolver
+
+ def __getattr__(self, name: str) -> Any:
+ if name == "map_index":
+ return self.resolver.public_map_index(self.value)
+ return getattr(self.value, name)
+
+
+def task_coordinate_response(
+ schema: type[Response], value: TaskCoordinate, resolver:
TaskCoordinateResolver
+) -> Response:
+ return schema.model_validate(TaskCoordinateView(value, resolver))
+
+
+def resolve_task_scope(
+ *,
+ dag_id: str,
+ run_id: str,
+ task_id: str,
+ session: Session,
+ dag_bag: DBDagBag,
+ map_index: int = -1,
+ region_id: UUID | None = None,
+ region_index: int | None = None,
+ all_map_indices: bool = False,
+) -> TaskScope:
+ """Resolve a public data address, including explicitly addressed retained
data."""
+ if region_index is not None and region_id is None:
+ raise HTTPException(status.HTTP_400_BAD_REQUEST, "region_index
requires region_id")
+ if region_id is not None and region_index is not None:
Review Comment:
When both region fields are given, the path or body `map_index` is dropped
without being compared, so `PATCH
.../taskInstances/m/7?region_id=R®ion_index=2` updates slot 2 and reports
`map_index: 2`. Could this return 400 when `map_index` isn't -1 and doesn't
match `region_index`?
##########
airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py:
##########
@@ -197,33 +221,107 @@ def _patch_ti_validate_request(
session: SessionDep,
map_index: int | None = -1,
update_mask: list[str] | None = None,
+ *,
+ lock: bool = True,
) -> tuple[SerializedDAG, list[TI], dict]:
+ _validate_region_selection(body)
dag = get_latest_version_of_dag(dag_bag, dag_id, session)
- if not dag.has_task(task_id):
+ if lock:
+ _lock_patch_runs(dag, dag_run_id, body, session)
+ if body.region_id is None and not dag.has_task(task_id):
raise HTTPException(status.HTTP_404_NOT_FOUND, f"Task '{task_id}' not
found in Dag '{dag_id}'")
query = (
select(TI)
.where(TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id ==
task_id)
.options(joinedload(TI.rendered_task_instance_fields))
)
- if map_index is not None:
- query = query.where(TI.map_index == map_index)
- else:
- query = query.order_by(TI.map_index)
-
+ if body.region_id is None:
+ resolver = TaskCoordinateResolver(dag_bag, session)
+ for version in session.scalars(
+ select(TI.dag_version_id)
+ .where(
+ TI.dag_id == dag_id,
+ TI.run_id == dag_run_id,
+ TI.task_id == task_id,
+ )
+ .distinct()
+ ):
+ if enclosing_loop(resolver.get_task(dag_id, dag_run_id, task_id,
dag_version_id=version)):
+ raise HTTPException(status.HTTP_409_CONFLICT, "Select a region
and index for this loop task")
+ scope = resolve_task_scope(
+ dag_id=dag_id,
+ run_id=dag_run_id,
+ task_id=task_id,
+ session=session,
+ dag_bag=dag_bag,
+ map_index=map_index if map_index is not None else -1,
+ region_id=body.region_id,
+ region_index=body.region_index,
+ all_map_indices=map_index is None,
+ )
+ query = query.where(TI.region_id == scope.region_id)
+ if body.region_id is not None or map_index is not None:
+ query = query.where(TI.region_index == scope.region_index)
+ query =
query.order_by(TI.region_index).execution_options(populate_existing=True)
+ if lock:
+ query = query.with_for_update(of=TI)
tis = session.scalars(query).all()
err_msg_404 = (
f"The Task Instance with dag_id: `{dag_id}`, run_id: `{dag_run_id}`,
task_id: `{task_id}` and map_index: `{map_index}` was not found",
)
if len(tis) == 0:
raise HTTPException(status.HTTP_404_NOT_FOUND, err_msg_404)
+ if body.region_id is not None and tis[0].dag_version_id is not None:
+ pinned_dag = dag_bag.get_dag(tis[0].dag_version_id, session=session)
+ if pinned_dag is not None:
+ dag = pinned_dag
data = _validate_patch_task_instance_body(body, update_mask)
return dag, list(tis), data
+def _validate_region_selection(body: PatchTaskInstanceBody) -> None:
+ if (body.region_id is None) != (body.region_index is None):
+ raise HTTPException(
+ status.HTTP_400_BAD_REQUEST, "region_id and region_index must be
supplied together"
+ )
+ if body.region_id is not None and (body.include_past or
body.include_future):
+ raise HTTPException(status.HTTP_400_BAD_REQUEST, "Regional selection
requires one explicit DagRun")
+
+
+def _lock_patch_runs(dag: SerializedDAG, run_id: str, body:
PatchTaskInstanceBody, session: Session) -> None:
+ run_ids = get_run_ids(dag, run_id, body.include_future, body.include_past,
session=session)
Review Comment:
`get_run_ids` costs up to four DagRun queries even when `include_past` and
`include_future` are both off, and this runs once per bulk entity (once per TI
in the "all map indexes" branch), each followed by a DagRun `FOR UPDATE` that
`_handle_regional_bulk` already holds. Marking a 1000-wide mapped task through
the bulk endpoint turns into thousands of extra statements while the run lock
is held, which also stalls restart acks for regional rows of that run. Could
this return `[run_id]` early when neither flag is set, and could the bulk path
skip the re-lock?
##########
airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py:
##########
@@ -3265,6 +3283,452 @@ def test_should_not_return_duplicate_runs(self,
test_client, session):
class TestPostClearTaskInstances(TestTaskInstanceEndpoint):
+ @pytest.mark.parametrize("one_run", [False, True])
+ def test_broad_clear_uses_pinned_loop_after_latest_definition_removes_it(
+ self, test_client, dag_maker, session, one_run
+ ):
+ @task_group
+ def body():
+ MockOperator(task_id="member")
+
+ with dag_maker(dag_id="changed_loop", serialized=True):
+ loop = create_loop(body, max_iterations=2)
+ dr = dag_maker.create_dagrun()
+ root = session.scalar(select(DynamicRegion).where(DynamicRegion.dag_id
== dr.dag_id))
+ for ti in dr.task_instances:
+ ti.state = State.SUCCESS
+ for loop_task in loop.iter_tasks():
+ session.add(
+ TaskInstance(
+ task=loop_task,
+ run_id=dr.run_id,
+ dag_version_id=dr.created_dag_version_id,
+ region_id=root.id,
+ region_index=1,
+ state=State.SUCCESS,
+ )
+ )
+ session.commit()
+ suffix_ids = {ti.id for ti in dr.get_task_instances(session=session)
if ti.region_index == 1}
+ dag_id, run_id, gate_id = dr.dag_id, dr.run_id, loop.gate_task_id
+ with dag_maker(dag_id=dag_id, serialized=True, session=session):
+ MockOperator(task_id="replacement")
+ session.commit()
+
+ response = test_client.post(
+ f"/dags/{dag_id}/clearTaskInstances",
+ json={
+ "task_ids": [gate_id],
+ "dry_run": False,
+ "only_failed": False,
+ **({"dag_run_id": run_id} if one_run else {}),
+ },
+ )
+
+ assert response.status_code == 200, response.text
+ session.expire_all()
+ assert len(suffix_ids) == 2
+ for identity in suffix_ids:
+ archived = session.get(TaskInstance, identity)
+ assert (archived.working_set, archived.archived_reason) == (None,
"superseded")
+ assert (
+ session.scalar(
+
select(func.count()).select_from(DynamicRegion).where(DynamicRegion.dag_id ==
dag_id)
+ )
+ == 2
+ )
+
+ @pytest.mark.parametrize("exact", [False, True])
+ def test_task_name_clear_refreshes_execution_after_waiting_for_dagrun_lock(
+ self, test_client, dag_maker, session, mocker, exact
+ ):
+ with dag_maker(serialized=True):
+ MockOperator(task_id="task")
+ dr = dag_maker.create_dagrun()
+ ti = dr.get_task_instance("task", session=session)
+ ti.state = State.SUCCESS
+ session.commit()
+ run_id, dag_id, original_id = dr.run_id, dr.dag_id, ti.id
+ execute = Session._execute_internal
+ bind = session.get_bind()
+ replacement_id = None
+
+ def overlap(request_session, statement, *args, **kwargs):
+ nonlocal replacement_id
+ run_lock = (
+ isinstance(statement, Select)
+ and statement._for_update_arg is not None
+ and any(getattr(table, "name", None) == "dag_run" for table in
statement.get_final_froms())
+ )
+ if run_lock and replacement_id is None:
+ replacement_id = original_id
+ with Session(bind=bind) as other:
+ successor = clear_task_instances([other.get(TaskInstance,
original_id)], session=other)[0]
+ successor.state = State.SUCCESS
+ other.commit()
+ replacement_id = successor.id
+ return execute(request_session, statement, *args, **kwargs)
+
+ mocker.patch.object(Session, "_execute_internal", autospec=True,
side_effect=overlap)
+ selection = {"task_instance_ids": [str(original_id)]} if exact else
{"task_ids": ["task"]}
+ response = test_client.post(
+ f"/dags/{dag_id}/clearTaskInstances",
+ json={"dag_run_id": run_id, **selection, "dry_run": False,
"only_failed": False},
+ )
+
+ assert response.status_code == (409 if exact else 200), response.text
+ if not exact:
+ assert response.json()["total_entries"] == 1
+ session.expire_all()
+ assert session.get(TaskInstance, original_id).working_set is None
+ assert (session.get(TaskInstance, replacement_id).working_set is True)
is exact
+
+ def test_exact_clear_rejects_missing_execution_without_changing_run(
+ self, test_client, dag_maker, session
+ ):
+ with dag_maker(serialized=True):
+ MockOperator(task_id="task")
+ dr = dag_maker.create_dagrun()
+ session.commit()
+ before = {ti.id for ti in dr.get_task_instances(session=session)}
+
+ response = test_client.post(
+ f"/dags/{dr.dag_id}/clearTaskInstances",
+ json={
+ "dag_run_id": dr.run_id,
+ "task_instance_ids": [str(uuid7())],
+ "dry_run": False,
+ "only_failed": False,
+ },
+ )
+
+ assert response.status_code == 404
+ session.expire_all()
+ assert dr.clear_number == 0
+ assert {ti.id for ti in dr.get_task_instances(session=session)} ==
before
+
+ @pytest.mark.parametrize(
+ "selection",
+ [
+ {"task_instance_ids": []},
+ {"task_instance_ids": ["00000000-0000-0000-0000-000000000001"]},
+ {
+ "task_instance_ids": ["00000000-0000-0000-0000-000000000001"],
+ "dag_run_id": "run",
+ "task_ids": ["task"],
+ },
+ {
+ "task_instance_ids": ["00000000-0000-0000-0000-000000000001"],
+ "dag_run_id": "run",
+ "task_group_id": "group",
+ },
+ {
+ "task_instance_ids": ["00000000-0000-0000-0000-000000000001"],
+ "dag_run_id": "run",
+ "include_future": True,
+ },
+ {
+ "task_instance_ids": ["00000000-0000-0000-0000-000000000001"],
+ "dag_run_id": "run",
+ "include_past": True,
+ },
+ {"dag_run_id": "run", "whole_expansion_ids":
["00000000-0000-0000-0000-000000000001"]},
+ ],
+ )
+ def test_exact_clear_rejects_conflicting_scope(self, test_client,
selection):
+ response =
test_client.post("/dags/example_python_operator/clearTaskInstances",
json=selection)
+ assert response.status_code == 422
+
+ @pytest.mark.parametrize("downstream", [False, True])
+ @pytest.mark.parametrize("later", [False, True])
+ @pytest.mark.parametrize("dry_run", [False, True])
+ @pytest.mark.parametrize("only_failed", [False, True])
+ def test_exact_loop_clear_keeps_iteration_scope(
+ self, test_client, dag_maker, session, downstream, later, dry_run,
only_failed
+ ):
+ @task_group
+ def body():
+ MockOperator(task_id="member")
+
+ with dag_maker(serialized=True):
+ loop = create_loop(body, max_iterations=3)
+ loop >> MockOperator(task_id="outside")
+ dr = dag_maker.create_dagrun()
+ root = session.scalar(
+ select(DynamicRegion).where(
+ DynamicRegion.dag_id == dr.dag_id, DynamicRegion.node_id ==
loop.group_id
+ )
+ )
+ for index in (1, 2):
+ for loop_task in loop.iter_tasks():
+ session.add(
+ TaskInstance(
+ task=loop_task,
+ run_id=dr.run_id,
+ dag_version_id=dr.created_dag_version_id,
+ region_id=root.id,
+ region_index=index,
+ state=State.SUCCESS,
+ )
+ )
+ session.commit()
+ tis = list(dr.get_task_instances(session=session))
+ seed = next(ti for ti in tis if ti.task_id == "body.member" and
ti.region_index == 1)
+ if only_failed:
+ seed.state = State.FAILED
+ session.commit()
+ before = {(ti.id, ti.state, ti.try_number) for ti in tis}
+ coordinates = {ti.id: (ti.task_id, str(ti.region_id), ti.region_index)
for ti in tis}
+
+ response = test_client.post(
+ f"/dags/{dr.dag_id}/clearTaskInstances",
+ json={
+ "dag_run_id": dr.run_id,
+ "dry_run": dry_run,
+ "task_instance_ids": [str(seed.id)],
+ "only_failed": only_failed,
+ "include_downstream": downstream,
+ "include_later_loop_iterations": later,
+ },
+ )
+
+ assert response.status_code == 200, response.text
+ expected = {seed.id}
+ if downstream and not only_failed:
+ expected.update(
+ ti.id
+ for ti in tis
+ if ti.task_id == "outside"
+ or (ti.region_index == 1 and ti.task_id == loop.gate_task_id)
+ or (later and ti.region_id == root.id and ti.region_index > 1)
+ )
+ assert {
+ (row["task_id"], row["region_id"], row["region_index"])
+ for row in response.json()["task_instances"]
+ } == {coordinates[value] for value in expected}
+ session.expire_all()
+ if dry_run:
+ assert {
+ (ti.id, ti.state, ti.try_number) for ti in
dr.get_task_instances(session=session)
+ } == before
+ else:
+ for identity, state, try_number in before:
+ if identity not in expected:
+ retained = session.get(TaskInstance, identity)
+ assert (retained.state, retained.try_number) == (state,
try_number)
+ elif state in (State.SUCCESS, State.FAILED):
+ historical = session.get(TaskInstance, identity)
+ assert historical.working_set is None
+ archived = coordinates[identity][2] > 1 and downstream and
later
+ assert historical.archived_reason == ("superseded" if
archived else "retry")
+
+ def
test_loop_clear_reports_conflict_when_a_worker_archived_the_execution_first(
+ self, test_client, dag_maker, session, mocker
+ ):
+ @task_group
+ def body():
+ MockOperator(task_id="member")
+
+ with dag_maker(serialized=True):
+ loop = create_loop(body, max_iterations=3)
+ dr = dag_maker.create_dagrun()
+ root = session.scalar(
+ select(DynamicRegion).where(
+ DynamicRegion.dag_id == dr.dag_id, DynamicRegion.node_id ==
loop.group_id
+ )
+ )
+ for loop_task in loop.iter_tasks():
+ session.add(
+ TaskInstance(
+ task=loop_task,
+ run_id=dr.run_id,
+ dag_version_id=dr.created_dag_version_id,
+ region_id=root.id,
+ region_index=1,
+ state=State.SUCCESS,
+ )
+ )
+ session.commit()
+ gate = next(
+ ti
+ for ti in dr.get_task_instances(session=session)
+ if ti.task_id == loop.gate_task_id and ti.region_index == 0
+ )
+ mocker.patch.object(
+ TaskInstance,
+ "archive",
+ autospec=True,
+ side_effect=ValueError("An archived task instance cannot be
archived again"),
+ )
+
+ response = test_client.post(
+ f"/dags/{dr.dag_id}/clearTaskInstances",
+ json={
+ "dag_run_id": dr.run_id,
+ "dry_run": False,
+ "only_failed": False,
+ "task_instance_ids": [str(gate.id)],
+ "include_later_loop_iterations": True,
+ },
+ )
+
+ assert response.status_code == 409, response.text
+
+ @pytest.mark.parametrize("downstream", [False, True])
+ @pytest.mark.parametrize("later", [False, True])
+ @pytest.mark.parametrize("seed_task", ["body.improve", "body.evaluate",
"gate"])
+ @pytest.mark.parametrize("exact", [False, True])
+ @pytest.mark.parametrize("only_failed", [False, True])
+ def test_loop_clear_dry_run_reports_what_the_real_clear_replaces(
+ self, test_client, dag_maker, session, downstream, later, seed_task,
exact, only_failed
+ ):
+ @task_group
+ def body():
+ improve = MockOperator(task_id="improve")
+ improve >> MockOperator(task_id="evaluate")
+
+ with dag_maker(serialized=True):
+ loop = create_loop(body, max_iterations=4)
+ loop >> MockOperator(task_id="finished")
+ dr = dag_maker.create_dagrun()
+ root = session.scalar(
+ select(DynamicRegion).where(
+ DynamicRegion.dag_id == dr.dag_id, DynamicRegion.node_id ==
loop.group_id
+ )
+ )
+ for index in (1, 2, 3):
+ for loop_task in loop.iter_tasks():
+ session.add(
+ TaskInstance(
+ task=loop_task,
+ run_id=dr.run_id,
+ dag_version_id=dr.created_dag_version_id,
+ region_id=root.id,
+ region_index=index,
+ state=State.SUCCESS,
+ )
+ )
+ for ti in dr.get_task_instances(session=session):
+ ti.state = State.FAILED if only_failed and ti.region_index in (1,
2) else State.SUCCESS
+ session.commit()
+ task_id = loop.gate_task_id if seed_task == "gate" else seed_task
+
+ def build_payload():
+ seed = session.scalar(
+ select(TaskInstance).where(
+ TaskInstance.run_id == dr.run_id,
+ TaskInstance.task_id == task_id,
+ TaskInstance.region_index == 1,
+ TaskInstance.working_set.is_(True),
+ )
+ )
+ return {
+ "dag_run_id": dr.run_id,
+ "only_failed": only_failed,
+ "include_downstream": downstream,
+ "include_later_loop_iterations": later,
+ **({"task_instance_ids": [str(seed.id)]} if exact else
{"task_ids": [task_id]}),
+ }
+
+ test_client.post(f"/dags/{dr.dag_id}/clearTaskInstances",
json={**build_payload(), "dry_run": False})
+ for ti in
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))):
+ ti.state = State.FAILED if only_failed and ti.region_index in (1,
2) else State.SUCCESS
+ session.commit()
+ payload = build_payload()
+ archived_before = set(
+ session.scalars(
+ select(TaskInstance.id)
+ .where(TaskInstance.working_set.is_(None))
+ .execution_options(include_all_attempts=True)
+ )
+ )
+ dry = test_client.post(f"/dags/{dr.dag_id}/clearTaskInstances",
json={**payload, "dry_run": True})
+ real = test_client.post(f"/dags/{dr.dag_id}/clearTaskInstances",
json={**payload, "dry_run": False})
+
+ assert dry.status_code == 200, dry.text
+ assert real.status_code == 200, real.text
+ session.expire_all()
+ archived_ids = {
+ ti.id
+ for ti in session.scalars(
+ select(TaskInstance)
+ .where(
+ TaskInstance.dag_id == dr.dag_id,
+ TaskInstance.working_set.is_(None),
+ )
+ .execution_options(include_all_attempts=True)
+ )
+ } - archived_before
+ reported = {
+ (row["task_id"], row["region_id"], row["region_index"]) for row in
dry.json()["task_instances"]
+ }
+ replaced = {
+ (ti.task_id, str(ti.region_id), ti.region_index)
+ for ti in
session.scalars(select(TaskInstance).where(TaskInstance.id.in_(archived_ids)))
Review Comment:
This query has no `include_all_attempts`, and an `IN` on ids isn't exempt
from the working-set filter (only `pk_` binds are), so every id in
`archived_ids` is filtered out and `replaced` is always empty. That makes
`reported == replaced` fail whenever the dry run reports anything, whatever the
route does. Adding `.execution_options(include_all_attempts=True)` here, or
collecting `(task_id, region_id, region_index)` in the `archived_ids` query
above, would let it check the parity it's named for.
##########
airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_dag_run.py:
##########
@@ -2255,6 +2255,40 @@ def test_clear_dag_run_dry_run(self, test_client,
session, body, dag_run_id, exp
)
assert logs == 0
+ @pytest.mark.usefixtures("configure_git_connection_for_dag_bundle")
+ def test_clear_dag_run_dry_run_reports_loop_coordinates(self, test_client,
dag_maker, session):
+ @task_group
+ def body():
+ EmptyOperator(task_id="work")
+
+ with dag_maker("dry_run_loop", serialized=True):
+ loop = create_loop(body, max_iterations=3)
+ run = dag_maker.create_dagrun(run_id="dry_run_loop_run")
+ region =
session.scalar(select(DynamicRegion).where(DynamicRegion.node_id ==
loop.group_id))
+ later = TaskInstance(
+ task=dag_maker.dag.get_task("body.work"),
+ run_id=run.run_id,
+ dag_version_id=run.created_dag_version_id,
+ region_id=region.id,
+ region_index=1,
+ )
+ later.state = State.SUCCESS
+ session.add(later)
+ session.commit()
+
+ response = test_client.post(
+ f"/dags/{run.dag_id}/dagRuns/{run.run_id}/clear", json={"dry_run":
True, "only_failed": False}
+ )
+
+ assert response.status_code == 200, response.text
+ work = [row for row in response.json()["task_instances"] if
row["task_id"] == "body.work"]
+ assert {row["region_index"] for row in work} == {0, 1}
+ assert {row["map_index"] for row in work} == {-1}
+ assert all(
+ row["loop_iterations"] == [{"loop_id": loop.group_id, "iteration":
row["region_index"]}]
Review Comment:
`loop_iterations` isn't a field on `TaskInstanceResponse` at this commit
(it's added in #74350), so `row["loop_iterations"]` raises `KeyError` here.
Could this assertion move over with the field?
##########
airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py:
##########
@@ -288,7 +446,23 @@ def _patch_task_instance_state(
task_instance_body: BulkTaskInstanceBody | PatchTaskInstanceBody,
data: dict,
session: Session,
+ selected: list[TI] | None = None,
+ commit: bool = True,
) -> list[TI]:
+ if task_instance_body.region_id is not None:
+ if selected is None:
Review Comment:
Every caller that sets `region_id` already passes `selected` (the routes and
`_handle_regional_bulk`), and `_perform_update` only ever sees non-regional
entities, so this fallback query can't run. Making `selected` required on the
regional branch, or dropping the fallback here and in
`_patch_task_group_state`, would make that explicit.
##########
airflow-core/src/airflow/models/loop_clear.py:
##########
@@ -0,0 +1,449 @@
+# 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
+
+from collections.abc import Collection, Iterator
+from dataclasses import dataclass
+from typing import TYPE_CHECKING, Literal
+from uuid import UUID
+
+from sqlalchemy import exists, select, tuple_
+
+from airflow.exceptions import AirflowClearRunningTaskException
+from airflow.models.dagbag import DBDagBag
+from airflow.models.dynamic_region import (
+ SENTINEL_REGION_ID,
+ DynamicRegion,
+ load_region_ancestry,
+ loop_position,
+)
+from airflow.models.task_coordinates import TaskCoordinateResolver,
enclosing_loop
+from airflow.models.taskinstance import TaskInstance,
_get_relevant_map_indexes, clear_task_instances
+from airflow.serialization.definitions.mappedoperator import
get_mapped_ti_count
+from airflow.utils.state import DagRunState, TaskInstanceState
+
+if TYPE_CHECKING:
+ from sqlalchemy.orm import Session
+
+ from airflow.serialization.definitions.taskgroup import
SerializedLoopTaskGroup
+
+
+@dataclass(frozen=True)
+class LoopClearScope:
+ """Live executions to retry and generated executions to archive."""
+
+ retry_ids: frozenset[UUID]
+ archive_ids: frozenset[UUID]
+
+
+def select_loop_clear_scope(
+ selected: Collection[TaskInstance],
+ *,
+ whole_expansion_ids: Collection[UUID] = (),
+ upstream: bool = False,
+ downstream: bool = True,
+ later_loop_iterations: bool = True,
+ session: Session,
+) -> LoopClearScope:
+ """Select loop clear executions; mutation callers must hold the DagRun
lock."""
+ if not selected:
+ return LoopClearScope(frozenset(), frozenset())
+ runs = {(ti.dag_id, ti.run_id) for ti in selected}
+ if len(runs) != 1:
+ raise ValueError("Loop clear selection must belong to one DagRun")
+ dag_id, run_id = runs.pop()
+ live = {
+ ti.id: ti
+ for ti in session.scalars(
+ select(TaskInstance)
+ .where(
+ TaskInstance.dag_id == dag_id,
+ TaskInstance.run_id == run_id,
+ TaskInstance.working_set.is_(True),
+ )
+ .execution_options(populate_existing=True)
+ )
+ }
+ selected_ids = {ti.id for ti in selected}
+ if not selected_ids <= live.keys():
+ raise ValueError("Loop clear selection contains executions that are no
longer live")
+ if not set(whole_expansion_ids) <= selected_ids:
+ raise ValueError("Whole-task selection requires an explicitly selected
execution")
+ resolver = TaskCoordinateResolver(DBDagBag(), session)
+ regions = load_region_ancestry(
+ {ti.region_id for ti in live.values()}, dag_id=dag_id, run_id=run_id,
session=session
+ )
+ retry = {ti_id: live[ti_id] for ti_id in selected_ids}
+ for ti in tuple(retry.values()):
+ task = resolver.get_task(dag_id, run_id, ti.task_id,
dag_version_id=ti.dag_version_id)
+ group = enclosing_loop(task)
+ if group is not None and loop_position(regions, ti.region_id,
ti.region_index, group.node_id) is None:
+ raise ValueError("Loop clear selection requires an execution
inside its pinned loop")
+ if ti.id in whole_expansion_ids:
+ if not task.get_needs_expansion():
+ raise ValueError("Whole-expansion selection requires a mapped
task")
+ retry.update(
+ (other.id, other)
+ for other in resolver.resolve(dag_id=dag_id, run_id=run_id,
task_id=ti.task_id, caller=ti)
+ )
+ for is_upstream, enabled in ((True, upstream), (False, downstream)):
+ if not enabled:
+ continue
+ for ti in tuple(retry.values()):
+ task = resolver.get_task(dag_id, run_id, ti.task_id,
dag_version_id=ti.dag_version_id)
+ contexts = resolver.producer_contexts(ti)
+ count = (
+ get_mapped_ti_count(task, run_id, session=session,
producer_contexts=contexts)
+ if task.get_needs_expansion() and ti.region_index >= 0
+ else None
+ )
+ for relative in task.get_flat_relatives(upstream=is_upstream):
+ indexes = _get_relevant_map_indexes(
+ task=task,
+ run_id=run_id,
+ map_index=resolver.public_map_index(ti),
+ relative=relative,
+ ti_count=count,
+ session=session,
+ producer_contexts=contexts,
+ )
+ relative_loop = enclosing_loop(relative)
+ task_loop = enclosing_loop(task)
+ matches: Collection[TaskInstance]
+ if relative_loop is not None and (
+ task_loop is None or task_loop.group_id !=
relative_loop.group_id
+ ):
+ matches = [
+ other
+ for other in live.values()
+ if other.task_id == relative.task_id
+ and (
+ indexes is None
+ or (isinstance(indexes, int) and
resolver.public_map_index(other) == indexes)
+ or (isinstance(indexes, range) and
resolver.public_map_index(other) in indexes)
+ )
+ ]
+ else:
+ matches = resolver.resolve(
+ dag_id=dag_id,
+ run_id=run_id,
+ task_id=relative.task_id,
+ caller=ti,
+ map_indexes=indexes,
+ )
+ retry.update((other.id, other) for other in matches)
+ archived: set[UUID] = set()
+ if later_loop_iterations:
+ for ti in retry.values():
+ task = resolver.get_task(dag_id, run_id, ti.task_id,
dag_version_id=ti.dag_version_id)
+ group = enclosing_loop(task)
+ if group is None or ti.task_id != group.gate_task_id:
+ continue
+ position = loop_position(regions, ti.region_id, ti.region_index,
group.node_id)
+ if position is None:
+ raise ValueError("Gate coordinates do not belong to its pinned
loop")
+ for other in live.values():
+ other_position = loop_position(regions, other.region_id,
other.region_index, group.node_id)
+ if (
+ other_position is not None
+ and other_position[0] == position[0]
+ and other_position[1] > position[1]
+ ):
+ archived.add(other.id)
+ return LoopClearScope(frozenset(retry.keys() - archived),
frozenset(archived))
+
+
+def _regions_for_run(ti: TaskInstance, session: Session) -> dict[UUID,
DynamicRegion]:
+ return {
+ region.id: region
+ for region in session.scalars(
+ select(DynamicRegion).where(DynamicRegion.dag_id == ti.dag_id,
DynamicRegion.run_id == ti.run_id)
+ )
+ }
+
+
+def _physical_loop_coordinate(
+ region_id: UUID, index: int, node_id: str, regions: dict[UUID,
DynamicRegion]
+) -> tuple[UUID, int] | None:
+ while region_id in regions:
+ region = regions[region_id]
+ if region.node_id == node_id:
+ return region_id, index
+ if region.parent_region_id is None or region.parent_region_index is
None:
+ return None
+ region_id, index = region.parent_region_id, region.parent_region_index
+ return None
+
+
+def _forks_after(region_id: UUID, regions: dict[UUID, DynamicRegion]) ->
Iterator[DynamicRegion]:
+ successors = {region.forked_from_region_id: region for region in
regions.values()}
+ while region_id in successors:
+ successor = successors[region_id]
+ yield successor
+ region_id = successor.id
+
+
+def loop_execution_is_superseded(ti: TaskInstance, *, session: Session) ->
bool:
+ """Identify a archiving loop coordinate for termination acknowledgement."""
+ return loop_coordinate_is_superseded(ti.region_id, ti.region_index,
_regions_for_run(ti, session))
+
+
+def loop_coordinate_is_superseded(region_id: UUID, index: int, regions:
dict[UUID, DynamicRegion]) -> bool:
+ """Test whether a coordinate or its parent lies beyond a durable fork
cut."""
+ while region_id in regions:
+ if any(region.resumes_from_index <= index for region in
_forks_after(region_id, regions)):
+ return True
+ region = regions[region_id]
+ if region.parent_region_id is None or region.parent_region_index is
None:
+ break
+ region_id, index = region.parent_region_id, region.parent_region_index
+ return False
+
+
+def loop_gate_waits_for_archival(gate: TaskInstance, *, session: Session) ->
bool:
+ """Hold a rerun gate until superseded executions of its later passes
finish termination."""
+ if gate.operator != "LoopGateOperator" or gate.region_id ==
SENTINEL_REGION_ID:
+ return False
+ if not session.scalar(
+ select(
+ exists().where(
+ DynamicRegion.dag_id == gate.dag_id,
+ DynamicRegion.run_id == gate.run_id,
+ DynamicRegion.forked_from_region_id.is_not(None),
+ )
+ )
+ ):
+ return False
+ regions = _regions_for_run(gate, session)
+ gate_region = regions.get(gate.region_id)
+ if gate_region is None:
+ return False
+ node_id = gate_region.node_id
+ gate_position = loop_position(regions, gate.region_id, gate.region_index,
node_id)
+ if gate_position is None:
+ raise ValueError("Gate coordinates do not belong to its pinned loop")
+ for region_id, region_index in session.execute(
+ select(TaskInstance.region_id, TaskInstance.region_index).where(
+ TaskInstance.dag_id == gate.dag_id,
+ TaskInstance.run_id == gate.run_id,
+ TaskInstance.working_set.is_(True),
+ TaskInstance.state == TaskInstanceState.RESTARTING,
+ )
+ ):
+ position = loop_position(regions, region_id, region_index, node_id)
+ if position is None or position[0] != gate_position[0] or position[1]
<= gate_position[1]:
+ continue
+ coordinate = _physical_loop_coordinate(region_id, region_index,
node_id, regions)
+ if coordinate is not None and any(
+ region.resumes_from_index <= coordinate[1] for region in
_forks_after(coordinate[0], regions)
+ ):
+ return True
+ return False
+
+
+def loop_gate_has_later_pass(gate: TaskInstance, group:
SerializedLoopTaskGroup, *, session: Session) -> bool:
+ """Find a live gate of a later pass in the same loop family, including
passes in forked regions."""
+ regions = _regions_for_run(gate, session)
+ position = loop_position(regions, gate.region_id, gate.region_index,
group.node_id)
+ if position is None:
+ raise ValueError("Gate coordinates do not belong to its pinned loop")
+ for region_id, region_index in session.execute(
+ select(TaskInstance.region_id, TaskInstance.region_index).where(
+ TaskInstance.working_set.is_(True),
+ TaskInstance.dag_id == gate.dag_id,
+ TaskInstance.run_id == gate.run_id,
+ TaskInstance.task_id == gate.task_id,
+ TaskInstance.region_index > gate.region_index,
+ )
+ ):
+ other = loop_position(regions, region_id, region_index, group.node_id)
+ if other is not None and other[0] == position[0]:
+ return True
+ return False
+
+
+def loop_successor_region(gate: TaskInstance, *, session: Session) -> UUID:
+ """Choose the region for a newly generated successor pass."""
+ region_id = gate.region_id
+ for fork in _forks_after(gate.region_id, _regions_for_run(gate, session)):
+ if fork.resumes_from_index <= gate.region_index + 1:
+ region_id = fork.id
+ return region_id
+
+
+def clear_loop_task_instances(
+ selected: Collection[TaskInstance],
+ *,
+ whole_expansion_ids: Collection[UUID] = (),
+ upstream: bool = False,
+ downstream: bool = True,
+ later_loop_iterations: bool = True,
+ session: Session,
+ dag_run_state: DagRunState | Literal[False] = DagRunState.QUEUED,
+ run_on_latest_version: bool = False,
+ prevent_running_task: bool | None = None,
+ whole_task_keys: Collection[tuple[str, str, str]] = (),
+) -> LoopClearScope:
+ """Clear selected loop tries and archive their replaced generated
suffixes."""
+ from airflow.models.dagrun import DagRun
+
+ if not selected:
+ return LoopClearScope(frozenset(), frozenset())
+ run_keys = {(ti.dag_id, ti.run_id) for ti in selected}
+ if len(run_keys) != 1:
+ raise ValueError("Loop clear selection must belong to one DagRun")
+ dag_id, run_id = run_keys.pop()
+ session.scalar(select(DagRun).where(DagRun.dag_id == dag_id, DagRun.run_id
== run_id).with_for_update())
+ scope = select_loop_clear_scope(
+ selected,
+ whole_expansion_ids=whole_expansion_ids,
+ upstream=upstream,
+ downstream=downstream,
+ later_loop_iterations=later_loop_iterations,
+ session=session,
+ )
+ apply_loop_clear_scope(
+ scope,
+ later_loop_iterations=later_loop_iterations,
+ session=session,
+ dag_run_state=dag_run_state,
+ run_on_latest_version=run_on_latest_version,
+ prevent_running_task=prevent_running_task,
+ whole_task_keys=whole_task_keys,
+ )
+ return scope
+
+
+def clear_task_instances_for_runs(
+ tis: Collection[TaskInstance],
+ *,
+ session: Session,
+ dag_run_state: DagRunState | Literal[False] = DagRunState.QUEUED,
+ run_on_latest_version: bool = False,
+ prevent_running_task: bool | None = None,
+ whole_task_keys: Collection[tuple[str, str, str]] = (),
+ later_loop_iterations: bool = True,
+) -> None:
+ """
+ Clear exactly the given executions, archiving the later passes of any loop
gate among them.
+
+ Runs without a selected gate are retried in place as an ordinary clear.
+ """
+ from airflow.models.dagrun import DagRun
+
+ if not tis:
+ return
+ run_keys = sorted({(ti.dag_id, ti.run_id) for ti in tis})
+ session.scalars(
Review Comment:
This is the first of three identical DagRun `FOR UPDATE`s on this path:
`clear_loop_task_instances` locks the run again (line 310), and
`clear_task_instances` locks it a third time. The same sorted-lock block
appears in about eight places in this PR, and the lock order the commit message
depends on lives in each copy. A small `lock_dag_runs(session, run_keys)`
helper taken once at the outer entry point would keep the ordering in one place.
##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -1468,6 +1468,8 @@ def recalculate(self) -> _UnfinishedStates:
@provide_session
def task_instance_scheduling_decisions(self, *, session: Session =
NEW_SESSION) -> TISchedulingDecision:
tis = self.get_task_instances(session=session, state=State.task_states)
+ reconciled = self._reconcile_legacy_expansions(session=session)
Review Comment:
This puts a `dynamic_region` query with three correlated subqueries on every
scheduling pass of every run (the `test_dag.py` count goes from 5 to 6), though
it only finds work after a whole-task clear of a pre-region mapped expansion.
Could the placeholder TI be created where that state appears instead, in
`clear_task_instances` when the legacy rows archive straight away and in
`complete_restart` when the last superseded sentinel row archives, keeping the
call in `verify_integrity` only as a backstop?
--
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]