Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions airflow-core/docs/authoring-and-scheduling/loops.rst
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ iterations.
This example improves an estimate of the square root of two until the error
is small enough:

.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
:start-after: [START refine_estimate]
:end-before: [END refine_estimate]

Expand Down Expand Up @@ -98,15 +98,15 @@ For a fixed-count loop, omit ``until``. The definition ``accumulate.loop(max_ite
three iterations, carrying results between them. Reaching the cap completes
a fixed-count loop successfully. Its gate is named ``__loop_gate`` within the group.

.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
:start-after: [START fixed_loop]
:end-before: [END fixed_loop]

For a task-group function with arguments, supply them with ``.partial()`` before
calling ``.loop()``. Use ``.override()`` to configure the group, for example to
give another loop a different ``group_id``:

.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
:start-after: [START partial_override_loop]
:end-before: [END partial_override_loop]

Expand Down Expand Up @@ -258,7 +258,7 @@ created.
In this example, ``choose_items`` returns the values to map over: ``[1, 2]`` in the
first iteration, then each previous result plus one.

.. exampleinclude:: /authoring-and-scheduling/examples/example_task_loops.py
.. exampleinclude:: /../src/airflow/example_dags/example_task_loops.py
:start-after: [START mapped_loop]
:end-before: [END mapped_loop]

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -427,12 +427,24 @@ def safe_extract_from_orm(cls, data: Any) -> Any:
return values


class LoopContext(BaseModel):
"""Pinned loop definition and enclosing iteration for a task execution."""

node_id: str
index: int = Field(ge=0)
max_iterations: int = Field(gt=0)
terminal_task_id: str
terminal_is_mapped: bool


class TIRunContext(BaseModel):
"""Response schema for TaskInstance run context."""

dag_run: DagRun
"""DAG run information for the task instance."""

loop: LoopContext | None = None

task_reschedule_count: int = 0
"""How many times the task has been rescheduled."""

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@
from opentelemetry.trace import StatusCode
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
from pydantic import JsonValue, ValidationError
from sqlalchemy import and_, func, or_, tuple_, update
from sqlalchemy import and_, func, or_, tuple_, union, update
from sqlalchemy.engine import CursorResult
from sqlalchemy.exc import DataError, NoResultFound, SQLAlchemyError
from sqlalchemy.orm import Session, contains_eager, joinedload
Expand All @@ -57,6 +57,7 @@
from airflow.api_fastapi.execution_api.datamodels.taskinstance import (
DagRunNoteUpdatePayload,
InactiveAssetsResponse,
LoopContext,
PreviousTIResponse,
PrevSuccessfulDagRunResponse,
TaskBreadcrumbsResponse,
Expand Down Expand Up @@ -95,11 +96,12 @@
from airflow.models.base import ID_LEN
from airflow.models.dag import DagModel
from airflow.models.dagbag import DBDagBag
from airflow.models.dagrun import DagRun as DR
from airflow.models.dagrun import DagRun as DR, InvalidLoopDecision
from airflow.models.dynamic_region import SENTINEL_REGION_ID, AmbiguousProducerError
from airflow.models.hitl import HITLDetail
from airflow.models.log import Log
from airflow.models.task_coordinates import (
LOOP_GATE_OPERATOR,
TaskCoordinateResolver,
build_coordinate_filters,
get_public_region,
Expand All @@ -113,7 +115,7 @@
from airflow.state import get_state_backend
from airflow.triggers.base import TriggerEvent
from airflow.utils.sqlalchemy import get_dialect_name
from airflow.utils.state import DagRunState, TaskInstanceState, TerminalTIState
from airflow.utils.state import DagRunState, IntermediateTIState, TaskInstanceState, TerminalTIState

router = VersionedAPIRouter()

Expand Down Expand Up @@ -177,6 +179,7 @@ def ti_run(
TI.dag_id,
TI.run_id,
TI.task_id,
TI.region_id,
TI.region_index,
TI.try_number,
TI.max_tries,
Expand Down Expand Up @@ -336,6 +339,19 @@ def ti_run(
should_retry=_is_eligible_to_retry(previous_state, ti.try_number, ti.max_tries),
multi_team=conf.getboolean("core", "multi_team"),
)
resolver = TaskCoordinateResolver(dag_bag, session)
if loop_context := resolver.loop_context(ti):
group, index = loop_context
terminal = resolver.get_task(
ti.dag_id, ti.run_id, group.terminal_task_id, dag_version_id=ti.dag_version_id
)
context.loop = LoopContext(
node_id=group.node_id,
index=index,
max_iterations=group.max_iterations,
terminal_task_id=group.terminal_task_id,
terminal_is_mapped=terminal.get_needs_expansion(),
)

# Only set for lang-SDK (foreign-runtime) tasks with a captured TaskFlow arg
# spec; the route excludes unset fields, keeping regular responses lean.
Expand Down Expand Up @@ -435,6 +451,46 @@ def ti_update_state(
raise HTTPException(status_code=409, detail={"reason": "invalid_state"})
return Response(status_code=status.HTTP_204_NO_CONTENT)

loop_group = None
loop_gate = None
if isinstance(ti_patch_payload, (TISuccessStatePayload, TITerminalStatePayload)):
gate_run = session.execute(
select(TI.dag_id, TI.run_id).where(
TI.id == task_instance_id,
TI.working_set.is_(True),
TI.operator == LOOP_GATE_OPERATOR,
)
).one_or_none()
if gate_run is not None:
session.execute(
select(DR).where(DR.dag_id == gate_run.dag_id, DR.run_id == gate_run.run_id).with_for_update()
).scalar_one()
loop_gate = session.scalar(
select(TI)
.where(TI.id == task_instance_id, TI.working_set.is_(True))
.with_for_update(of=TI)
.execution_options(populate_existing=True)
)
if loop_gate is not None:
try:
loop_context = TaskCoordinateResolver(dag_bag, session).loop_context(loop_gate)
except TaskNotFound:
# A gate renamed mid-run on an unversioned bundle leaves this attempt without a loop.
log.warning("Loop gate is missing from its pinned Dag version; failing it")
if isinstance(ti_patch_payload, TISuccessStatePayload):
ti_patch_payload = _build_rejected_gate_payload(
ti_patch_payload,
reason="Loop gate is missing from its pinned Dag version",
retry=False,
)
loop_context = None
if loop_context is None:
loop_gate = None
elif loop_context[0].gate_task_id != loop_gate.task_id:
raise HTTPException(status_code=409, detail={"reason": "invalid_loop_gate"})
else:
loop_group = loop_context[0]

old = (
select(
TI.state,
Expand Down Expand Up @@ -522,6 +578,25 @@ def ti_update_state(
detail={"reason": "invalid_partition_key", "message": str(e)},
) from e

gate_completed = False
if (
loop_gate is not None
and loop_group is not None
and isinstance(ti_patch_payload, TISuccessStatePayload)
):
try:
loop_gate.dag_run.complete_loop_gate(
loop_gate, loop_group, TaskInstanceState.SUCCESS, session=session
)
gate_completed = True
except InvalidLoopDecision as error:
log.warning("Loop gate success rejected", error=str(error))
ti_patch_payload = _build_rejected_gate_payload(
Comment thread
ashb marked this conversation as resolved.
ti_patch_payload,
reason=f"Loop gate success rejected: {error}",
retry=_is_eligible_to_retry(previous_state, try_number, max_tries),
)

# We exclude_unset to avoid updating fields that are not set in the payload
data = ti_patch_payload.model_dump(
exclude={"task_outlets", "outlet_events", "retry_delay_seconds", "retry_reason"},
Expand All @@ -545,6 +620,8 @@ def ti_update_state(
# Let DataErrorHandler return a 422 instead of silently marking the TI FAILED below.
raise
except Exception:
if loop_gate is not None:
raise
Comment thread
ashb marked this conversation as resolved.
# Set a task to failed in case any unexpected exception happened during task state update
log.exception(
"Error updating Task Instance state. Setting the task to failed.",
Expand Down Expand Up @@ -595,6 +672,9 @@ def ti_update_state(
# Defer to app-level SQLAlchemyError handler (returns HTTP 500).
raise

if loop_gate is not None and loop_group is not None and not gate_completed:
loop_gate.dag_run.complete_loop_gate(loop_gate, loop_group, updated_state, session=session)

if updated_state == TaskInstanceState.SUCCESS:
if conf.getboolean("state_store", "clear_on_success"):
scope = TaskScope(
Expand Down Expand Up @@ -628,6 +708,15 @@ def ti_update_state(
callback()


def _build_rejected_gate_payload(
payload: TISuccessStatePayload, *, reason: str, retry: bool
) -> TIRetryStatePayload | TITerminalStatePayload:
carried = payload.model_dump(include={"end_date", "rendered_map_index"}, exclude_unset=True)
if retry:
return TIRetryStatePayload(state=IntermediateTIState.UP_FOR_RETRY, retry_reason=reason, **carried)
return TITerminalStatePayload(state=TerminalStateNonSuccess.FAILED, retry_reason=reason, **carried)


def _emit_task_span(ti, state, *, resolver: TaskCoordinateResolver):
# just to be safe
if not ti.dag_run:
Expand Down Expand Up @@ -941,29 +1030,27 @@ def ti_skip_downstream(
for task_id, index in targets
if not resolver.has_regions(dag_id, run_id, task_id)
]
selected_ids = {
ti.id
selects = [
resolver.select_skip_target_ids(caller=caller, task_id=task_id, map_indexes=index)
for task_id, index in targets
if (task_id, index) not in plain
for ti in resolver.resolve(
dag_id=dag_id, run_id=run_id, task_id=task_id, caller=caller, map_indexes=index
)
}
]
if plain:
selected_ids.update(
session.scalars(
resolver.select_legacy_task_ids(
dag_id=dag_id,
run_id=run_id,
task_ids=[task_id for task_id, index in plain if index is None],
slots=[(task_id, index) for task_id, index in plain if index is not None],
)
selects.append(
resolver.select_legacy_task_ids(
dag_id=dag_id,
run_id=run_id,
task_ids=[task_id for task_id, index in plain if index is None],
slots=[(task_id, index) for task_id, index in plain if index is not None],
)
)
except AmbiguousProducerError as error:
raise HTTPException(status.HTTP_409_CONFLICT, str(error)) from error
except ValueError as error:
raise HTTPException(status.HTTP_400_BAD_REQUEST, str(error)) from error
if not selects:
return
selected_ids = set(session.scalars(union(*selects)))

# Don't overwrite tasks that are already executing or finished.
# See: https://github.com/apache/airflow/issues/59378
Expand All @@ -981,6 +1068,7 @@ def ti_skip_downstream(
query = (
update(TI)
.where(
TI.working_set.is_(True),
TI.id.in_(selected_ids),
skippable_state_clause,
)
Expand Down
Loading
Loading