kaxil commented on code in PR #74345:
URL: https://github.com/apache/airflow/pull/74345#discussion_r4198200394
##########
task-sdk/src/airflow/sdk/execution_time/task_runner.py:
##########
@@ -249,6 +250,12 @@ class RuntimeTaskInstance(TaskInstance):
sentry_integration: str = ""
+ def model_post_init(self, context: Any) -> None:
Review Comment:
`region_id` and `region_index` reach `RuntimeTaskInstance` here, but the
`TaskScope` that `get_template_context` builds for `task_state_store` still
passes only `map_index`, which is `-1` for a non-mapped task inside a loop.
With `[workers] state_store_backend` set to the common.io object-storage
backend, pass 0 and pass 1 of a looped task both write
`<base>/<dag>/<run>/<task>/-1/<key>`, so pass 1 overwrites the object that pass
0's row still points to (and a `delete`/`clear` in pass 1 unlinks it). Should
that scope be built from the region coordinates, `TaskScope(...,
map_index=self.region_index, region_id=self.region_id)`, with a task-runner
test that checks the accessor's scope carries the TI's region?
##########
airflow-core/src/airflow/utils/log/task_log_address.py:
##########
@@ -0,0 +1,219 @@
+# 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, Mapping
+from functools import cache
+from pathlib import PurePosixPath
+from string import Formatter
+from typing import TYPE_CHECKING
+from uuid import UUID
+
+import attrs
+import jinja2
+from jinja2.meta import find_undeclared_variables
+from sqlalchemy import select, tuple_
+
+from airflow.models.dagrun import DagRun
+from airflow.models.dynamic_region import SENTINEL_REGION_ID, DynamicRegion
+from airflow.models.serialized_dag import SerializedDagModel
+from airflow.models.tasklog import LogTemplate
+from airflow.serialization.definitions.taskgroup import SerializedLoopTaskGroup
+from airflow.utils.helpers import render_template
+from airflow.utils.session import create_session
+
+if TYPE_CHECKING:
+ from sqlalchemy.orm import Session
+
+ from airflow.models.taskinstance import TaskInstance
+
+
[email protected](frozen=True)
+class TaskLogContext:
+ """Prepared, immutable inputs for query-free task log filename
rendering."""
+
+ filename_template: str
+ logical_date: str
+ data_interval_start: str
+ data_interval_end: str
+ log_position: str = ""
+ map_index: int = -1
+
+
+def region_log_position(
+ region_id: UUID,
+ region_index: int,
+ *,
+ regions: Mapping[UUID, DynamicRegion],
+ node_kinds: Mapping[str, str],
+) -> str:
+ """Render the loop nesting and fork position of one coordinate, empty
outside any loop."""
+ frames = []
+ inside_loop = False
+ while region_id != SENTINEL_REGION_ID:
+ region = regions[region_id]
+ kind = node_kinds[region.node_id]
+ if kind not in {"loop", "map"}:
+ raise ValueError(f"Unsupported log-address region kind: {kind}")
+ position = 1
+ predecessor = region
+ while predecessor.forked_from_region_id is not None:
+ predecessor = regions[predecessor.forked_from_region_id]
+ position += 1
+ suffix = f".{position}" if position > 1 else ""
+ label = "pass" if kind == "loop" else "map"
+ inside_loop = inside_loop or kind == "loop"
+ frames.append(f"{label}={region_index}{suffix}")
+ if region.parent_region_id is None:
+ break
+ if TYPE_CHECKING:
+ assert region.parent_region_index is not None
+ region_id, region_index = region.parent_region_id,
region.parent_region_index
+ return "/".join(reversed(frames)) if inside_loop else ""
+
+
+class _LogTaskView:
+ def __init__(self, ti, map_index: int):
+ self._ti = ti
+ self.map_index = map_index
+
+ def __getattr__(self, name: str) -> object:
+ return getattr(self._ti, name)
+
+
+@cache
+def _compile_template(template: str) -> tuple[jinja2.Template | str, bool]:
+ if "{{" in template:
+ env = jinja2.Environment(autoescape=False) # nosec B701: Render a
filesystem path, not HTML.
+ return env.from_string(template), "log_position" in
find_undeclared_variables(env.parse(template))
+ fields = {field for _, field, _, _ in Formatter().parse(template)}
+ return template, "log_position" in fields
+
+
+def render_task_log_filename(ti, try_number: int, *, context: TaskLogContext)
-> str:
+ """Render a pinned template and apply its regional address exactly once."""
+ template, consumes_position = _compile_template(context.filename_template)
+ task_view = _LogTaskView(ti, context.map_index)
+ if isinstance(template, jinja2.Template):
+ filename = render_template(
+ template,
+ {
+ "ti": task_view,
+ "ts": context.logical_date,
+ "try_number": try_number,
+ "log_position": context.log_position,
+ },
+ native=False,
+ )
+ else:
+ filename = template.format(
+ dag_id=ti.dag_id,
+ task_id=ti.task_id,
+ run_id=ti.run_id,
+ data_interval_start=context.data_interval_start,
+ data_interval_end=context.data_interval_end,
+ logical_date=context.logical_date,
+ try_number=try_number,
+ log_position=context.log_position,
+ )
+ if context.log_position and not consumes_position:
+ path = PurePosixPath(filename)
+ return str(path.parent / context.log_position / path.name)
+ return filename
+
+
+def prepare_task_log_contexts(
+ tis: Collection[TaskInstance],
+ *,
+ session: Session | None = None,
+) -> dict[UUID, TaskLogContext]:
+ """Load shared address inputs in batches before rendering workloads or
reading logs."""
+ if not tis:
+ return {}
+ if session is None:
+ with create_session(scoped=False) as session:
+ return prepare_task_log_contexts(tis, session=session)
+ run_keys = {(ti.dag_id, ti.run_id) for ti in tis}
+ runs = {
+ (run.dag_id, run.run_id): run
+ for run in session.scalars(select(DagRun).where(tuple_(DagRun.dag_id,
DagRun.run_id).in_(run_keys)))
+ }
+ template_ids = {run.log_template_id for run in runs.values() if
run.log_template_id is not None}
+ templates: dict[int | None, str | None] = {
+ template.id: template.filename
+ for template in
session.scalars(select(LogTemplate).where(LogTemplate.id.in_(template_ids)))
+ }
+ if any(run.log_template_id is None for run in runs.values()):
+ templates[None] =
session.scalar(select(LogTemplate.filename).order_by(LogTemplate.id).limit(1))
+ regional = [ti for ti in tis if ti.region_id != SENTINEL_REGION_ID]
+ regions: dict[UUID, DynamicRegion] = {}
+ dags = {}
+ node_kinds = {}
+ if regional:
+ frontier = {ti.region_id for ti in regional}
+ while frontier:
+ rows =
session.scalars(select(DynamicRegion).where(DynamicRegion.id.in_(frontier))).all()
+ regions.update((row.id, row) for row in rows)
+ frontier = {
+ ancestor_id
+ for row in rows
+ for ancestor_id in (row.parent_region_id,
row.forked_from_region_id)
+ if ancestor_id is not None and ancestor_id not in regions
+ }
+ version_ids = {
+ ti.dag_version_id or runs[ti.dag_id,
ti.run_id].created_dag_version_id for ti in regional
+ }
+ for row in session.scalars(
+
select(SerializedDagModel).where(SerializedDagModel.dag_version_id.in_(version_ids))
+ ):
+ row.load_op_links = False
+ dag = row.dag
+ dags[row.dag_version_id] = dag
+ node_kinds[row.dag_version_id] = {
+ group_id: "loop"
+ for group_id, group in
dag.task_group.get_task_group_dict().items()
+ if group_id is not None and isinstance(group,
SerializedLoopTaskGroup)
+ } | {task.task_id: "map" for task in dag.tasks if
task.get_needs_expansion()}
+ contexts = {}
+ for ti in tis:
+ run = runs[ti.dag_id, ti.run_id]
+ template = templates.get(run.log_template_id)
+ if template is None:
+ raise ValueError(
+ f"No log_template entry found for ID {run.log_template_id!r}. "
+ "Please make sure you set up the metadatabase correctly."
+ )
+ position = ""
+ map_index = ti.map_index
+ if ti.region_id != SENTINEL_REGION_ID:
+ version_id = ti.dag_version_id or run.created_dag_version_id
+ if version_id is None:
+ raise ValueError(f"Regional task instance {ti.id} has no
pinned DAG version")
+ dag = dags[version_id]
+ position = region_log_position(
+ ti.region_id, ti.region_index, regions=regions,
node_kinds=node_kinds[version_id]
+ )
+ map_index = ti.region_index if
dag.get_task(ti.task_id).get_needs_expansion() else -1
Review Comment:
This copy of the public `map_index` rule has no fallback for a TI whose
pinned definition no longer has the task. `dag.get_task(ti.task_id)` raises
`TaskNotFound` here, and `node_kinds[region.node_id]` in `region_log_position`
raises `KeyError` when the loop group or the mapping is gone, while
`TaskCoordinateResolver.public_map_index` handles the same case by falling back
to `_is_mapped_region`. The scheduler and triggerer prepare one batch for all
their workloads, so a single such row (for example a mapped task removed or
unmapped mid-run on an unversioned bundle, once mapped expansions carry
regions, whose unfinished TIs get re-pinned to the new version) fails the whole
enqueue with nothing catching it. The triggerer used to build a plain workload
for exactly this case. Could this fall back to the stored region ("map" when
`region.node_id == ti.task_id`, otherwise "loop") or reuse the resolver,
instead of raising?
##########
go-sdk/airflow/spec.gen.go:
##########
@@ -84,6 +98,9 @@ type TaskGroupSpec struct {
// GroupDisplayName corresponds to the JSON schema field
"group_display_name".
GroupDisplayName string
+ // Loop corresponds to the JSON schema field "loop".
+ Loop TaskGroupLoop
Review Comment:
Adding `loop` to the task group schema regenerated `TaskGroupLoop` and
`TaskGroupSpec.Loop` into the public Go package, but nothing in the Go SDK
reads or serializes them, and the gate and terminal ids are derived by the
Python loop builder rather than set by an author. It also fails
`TestShapeForAuthoringShapesTheCoreSchema` in `go-sdk/internal/genspec`, whose
pinned TaskGroupSpec property list now gains `loop`. Adding `loop` to
`taskGroupShape.exclude` in `authoring.go` (next to `is_mapped`) and
regenerating should fix both.
##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py:
##########
@@ -119,21 +134,103 @@ def has_xcom_access(
log = logging.getLogger(__name__)
-async def xcom_query(
+def _build_xcom_read(
+ *,
dag_id: str,
run_id: str,
task_id: str,
key: str,
+ session: Session,
+ dag_bag: DBDagBag,
+ token: TIToken,
+ map_index: int | None = None,
+ region_id: UUID | None = None,
+ region_index: int | None = None,
+ include_prior_dates: bool = False,
+) -> Select:
+ """Select the XCom rows of the producers visible to the calling task
instance."""
+ resolver = TaskCoordinateResolver(dag_bag, session)
+ read = partial(
+ XComModel.get_many,
+ dag_ids=dag_id,
+ run_id=run_id,
+ task_ids=task_id,
+ key=key,
+ include_prior_dates=include_prior_dates,
+ )
+ try:
+ if region_index is not None and region_id is None:
+ raise ValueError("region_index requires region_id")
+ if region_id is not None and region_index is not None:
+ return read(region_id=region_id, map_indexes=region_index)
+ if region_id == SENTINEL_REGION_ID or not resolver.has_regions(
+ dag_id, None if include_prior_dates else run_id, task_id
+ ):
+ return read(region_id=SENTINEL_REGION_ID, map_indexes=map_index)
+ caller = session.get(TaskInstance, token.id)
+ if include_prior_dates:
+ if region_id is not None:
+ raise ValueError(
+ "Prior-run lookup requires producer coordinates resolved
separately for each run"
+ )
+ try:
+ task = resolver.get_task(
+ dag_id,
+ run_id,
+ task_id,
+ dag_version_id=caller.dag_version_id
+ if caller and (caller.dag_id, caller.run_id) == (dag_id,
run_id)
+ else None,
+ )
+ except TaskNotFound:
Review Comment:
For a cross-Dag `xcom_pull(..., include_prior_dates=True)` the `run_id` is
the caller's run, which often doesn't exist in the producer Dag.
`_producer_task` then raises `ValueError("Pinned Dag for ... not found")`,
which this `except TaskNotFound` lets through to the 400 handler, so the task
fails where the same pull used to get a 404 and return None. Today that needs a
loop producer, but once mapped expansions get regions `has_regions` is true for
any expanded mapped task. Could this branch skip the loop check and fall
through to the plain read when the target run doesn't exist? Catching
`ValueError` here would also swallow `AmbiguousProducerError`, since it
subclasses `ValueError`.
##########
airflow-core/src/airflow/executors/workloads/task.py:
##########
@@ -100,11 +102,14 @@ def make(
generator: JWTGenerator | None = None,
bundle_info: BundleInfo | None = None,
sentry_integration: str = "",
+ log_context: TaskLogContext | None = None,
) -> ExecuteTask:
"""Create an ExecuteTask workload from a TaskInstance ORM model."""
- from airflow.utils.helpers import log_filename_template_renderer
+ from airflow.utils.log.task_log_address import
prepare_task_log_contexts, render_task_log_filename
- ser_ti = TaskInstanceDTO.model_validate(ti, from_attributes=True)
+ if log_context is None:
+ log_context = prepare_task_log_contexts([ti],
session=object_session(ti))[ti.id]
+ ser_ti = task_instance_to_runtime(ti, model=TaskInstanceDTO,
map_index=log_context.map_index)
Review Comment:
With this, `workload.ti.map_index` (and so `TaskInstanceDTO.key`) is -1 for
every loop pass, while the ORM `TaskInstance.map_index` and `.key` still carry
the stored pass index, so one attempt has two identities depending on who
builds the key. Several consumers compare across the two:
`BaseExecutor.has_task` compares the ORM key with the registered DTO key, so it
is always False for a loop pass on executors without UUID keys; the K8s
executor's `revoke_task` and live-log selectors add `map_index=<k>` but the pod
was labelled from the workload with no `map_index`; and the Elasticsearch
`log_id` is rendered with -1 by the writer and with k by the API server. Could
executor and storage identity stay on `region_id`/`region_index`, with only the
user-facing value using the public index?
##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -1268,15 +1306,18 @@ def get_task_instance_count(
logical_dates: Annotated[list[UtcDateTime] | None, Query()] = None,
run_ids: Annotated[list[str] | None, Query()] = None,
states: Annotated[list[str] | None, Query()] = None,
+ region_id: UUID | None = None,
+ region_index: int | None = None,
) -> int:
"""Get the count of task instances matching the given criteria."""
query = select(func.count()).select_from(TI).where(TI.dag_id == dag_id)
if task_ids:
query = query.where(TI.task_id.in_(task_ids))
- if map_index is not None:
- query = query.where(TI.map_index == map_index)
+ query = _filter_task_coordinates(
Review Comment:
`/states` and `/previous` now return 409 when several live TIs share a
public coordinate, but `/count` counts each loop pass. So
`ExternalTaskSensor(external_task_ids=["loop.body"])` gets 3 for a loop that
ran three passes, `3 / 1 != 1`, and it pokes until timeout, while the same
sensor with `external_task_group_id` fails straight away on the 409. Should
`/count` apply the same check (group by run, task and public index, 409 when a
group has more than one live row)?
##########
airflow-core/src/airflow/utils/log/file_task_handler.py:
##########
@@ -516,41 +515,12 @@ def close(self):
if self.handler:
self.handler.close()
- @provide_session
- def _render_filename(self, ti: TaskInstance, try_number: int, *,
session=NEW_SESSION) -> str:
+ def _render_filename(self, ti: TaskInstance, try_number: int, *, session:
Session | None = None) -> str:
"""Return the worker log filename."""
- dag_run = ti.get_dagrun(session=session)
-
- date = dag_run.logical_date or dag_run.run_after
- formatted_date = date.isoformat()
-
- template = dag_run.get_log_template(session=session).filename
- str_tpl, jinja_tpl = parse_template_string(template)
- if jinja_tpl:
- return render_template(
- jinja_tpl, {"ti": ti, "ts": formatted_date, "try_number":
try_number}, native=False
- )
-
- if str_tpl:
- data_interval = (dag_run.data_interval_start,
dag_run.data_interval_end)
- if data_interval[0]:
- data_interval_start = data_interval[0].isoformat()
- else:
- data_interval_start = ""
- if data_interval[1]:
- data_interval_end = data_interval[1].isoformat()
- else:
- data_interval_end = ""
- return str_tpl.format(
- dag_id=ti.dag_id,
- task_id=ti.task_id,
- run_id=ti.run_id,
- data_interval_start=data_interval_start,
- data_interval_end=data_interval_end,
- logical_date=formatted_date,
- try_number=try_number,
- )
- raise RuntimeError(f"Unable to render log filename for {ti}. This
should never happen")
+ from airflow.utils.log.task_log_address import
prepare_task_log_contexts, render_task_log_filename
+
+ context = prepare_task_log_contexts([ti], session=session)[ti.id]
Review Comment:
The edge3 log push still resolves the TI with
`TaskInstance.get_task_instance(dag_id, run_id, task_id, map_index)` before
calling this. That lookup only matches the sentinel region, so for a loop pass
it returns None and `prepare_task_log_contexts([None])` raises
`AttributeError`, which turns every log push for a loop task on an Edge worker
into a 500. Could edge3 look the TI up by `EdgeJobModel.task_instance_id`, or
use the workload's already rendered log path, so it gets the same address as
the worker? Its `@cache` keyed by `TaskInstanceKey` would also map all passes
to one path.
##########
task-sdk/src/airflow/sdk/execution_time/schema/versions/v2026_10_30.py:
##########
@@ -29,8 +29,60 @@
from cadwyn import VersionChange, schema
from airflow.dag_processing.processor import DagFileParsingResult # noqa:
SDK002
-from airflow.sdk.api.datamodels._generated import TIRunContext
-from airflow.sdk.execution_time.comms import TaskState
+from airflow.sdk.api.datamodels._generated import PreviousTIResponse,
TaskInstance, TIRunContext
+from airflow.sdk.execution_time.comms import (
+ DeleteXCom,
+ GetPreviousTI,
+ GetTaskBreadcrumbs,
+ GetTaskStates,
+ GetTICount,
+ GetXCom,
+ GetXComCount,
+ GetXComSequenceItem,
+ GetXComSequenceSlice,
+ SetXCom,
+ TaskState,
+)
+
+
+class AddRegionSelectors(VersionChange):
Review Comment:
Is anything going to send these selectors? Loops resolve producers from the
caller's token on the server, and through the top of the stack no SDK path
(`xcom_pull`, `XCom.set`, `get_ti_count`, `get_task_states`, `get_previous_ti`)
passes `region_id`/`region_index`, yet the `/previous` 409 and the prior-date
XCom 400 tell the caller to supply explicit coordinates. If there is no caller
yet, keeping only the response-side coordinate fields and adding request
selectors with their first user would avoid freezing unused surface into
2026-10-30.
##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -893,28 +904,43 @@ def ti_skip_downstream(
task_instance_id: UUID,
ti_patch_payload: TISkippedDownstreamTasksStatePayload,
session: SessionDep,
+ dag_bag: DagBagDep,
):
bind_contextvars(ti_id=str(task_instance_id))
log.info("Skipping downstream tasks",
task_count=len(ti_patch_payload.tasks))
now = timezone.utcnow()
tasks = ti_patch_payload.tasks
- query_result = session.execute(select(TI.dag_id, TI.run_id).where(TI.id ==
task_instance_id))
- row_result = query_result.fetchone()
- if row_result is None:
+ caller = session.get(TI, task_instance_id)
+ if caller is None or caller.working_set is not True:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={"reason": "not_found", "message": "Task Instance not
found"},
)
- dag_id, run_id = row_result
+ dag_id, run_id = caller.dag_id, caller.run_id
log.debug("Retrieved DAG and run info", dag_id=dag_id, run_id=run_id)
- # A bare task_id skips every TI of that task, so an already expanded
mapped task
- # (e.g. one mapped over a literal list) is skipped too, not only map_index
-1.
- task_ids = [task for task in tasks if isinstance(task, str)]
- ti_keys = [task for task in tasks if isinstance(task, tuple)]
- log.debug("Prepared task IDs for skipping", task_ids=task_ids,
ti_keys=ti_keys)
+ log.debug("Prepared task IDs for skipping", tasks=tasks)
+ resolver = TaskCoordinateResolver(dag_bag, session)
+ try:
+ # A bare task_id skips every TI of that task, so an already expanded
mapped task
+ # (e.g. one mapped over a literal list) is skipped too, not only
map_index -1.
+ selected_ids = {
+ ti.id
+ for task in tasks
+ for ti in resolver.resolve(
Review Comment:
This runs `resolver.resolve` once per downstream task, and each call issues
a `has_regions` EXISTS plus a TI select, where the old code was one select and
one bulk UPDATE. A `ShortCircuitOperator` with
`ignore_downstream_trigger_rules=True` sends every downstream task, so a Dag
with a few hundred of them now does a few hundred round trips per skip request,
loops or not. Could the non-regional tasks keep the old `task_id IN (...)` /
`(task_id, map_index)` path, with one query to find which tasks have regions
and `resolve` only for those?
##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -1257,6 +1281,20 @@ async def get_previous_successful_dagrun(
return PrevSuccessfulDagRunResponse.model_validate(dag_run)
+def _filter_task_coordinates(
Review Comment:
This and `_coordinate_filters` in `xcoms.py` do the same validation and
filtering with slightly different rules (this one always applies `map_index`,
the XCom one only when `region_id` is None). Could one helper in
`models/task_coordinates.py` serve both, so the rules can't drift?
##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -170,6 +172,7 @@ def ti_run(
TI.run_id,
TI.task_id,
TI.map_index,
+ TI.region_id,
Review Comment:
Nothing in `ti_run` reads `region_id` from this row, so this column can go.
##########
airflow-core/src/airflow/models/task_coordinates.py:
##########
@@ -0,0 +1,279 @@
+# 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 typing import TYPE_CHECKING, Protocol
+from uuid import UUID
+
+import attrs
+from sqlalchemy import case, or_, select
+
+from airflow.exceptions import TaskNotFound
+from airflow.models.dagrun import DagRun
+from airflow.models.dynamic_region import (
+ SENTINEL_REGION_ID,
+ AmbiguousProducerError,
+ DynamicRegion,
+ ProducerContext,
+ resolve_current_producers,
+)
+from airflow.models.taskinstance import TaskInstance
+from airflow.serialization.definitions.taskgroup import SerializedLoopTaskGroup
+
+if TYPE_CHECKING:
+ from collections.abc import Collection
+
+ from sqlalchemy.orm import Session
+ from sqlalchemy.sql.elements import ColumnElement
+
+ from airflow.models.dagbag import DBDagBag
+ from airflow.serialization.definitions.dag import SerializedDAG,
SerializedOperator
+
+
+class TaskCoordinate(Protocol):
+ """Stored task identity shared by live and retained task data."""
+
+ dag_id: str
+ run_id: str
+ task_id: str
+ region_id: UUID
+ region_index: int
+
+
+def enclosing_loop(task: SerializedOperator) -> SerializedLoopTaskGroup | None:
+ group = task.task_group
+ while group is not None:
+ if isinstance(group, SerializedLoopTaskGroup):
+ return group
+ group = group.parent_group
+ return None
+
+
+def public_map_index_expression(model) -> ColumnElement[int]:
+ mapped_region = (
+ select(DynamicRegion.id)
+ .where(DynamicRegion.id == model.region_id, DynamicRegion.node_id ==
model.task_id)
+ .correlate(model)
+ .exists()
+ )
+ return case(
+ (or_(model.region_id == SENTINEL_REGION_ID, mapped_region),
model.region_index),
+ else_=-1,
+ )
+
+
[email protected]
+class TaskCoordinateResolver:
+ """Resolve task coordinates against the definitions pinned to their
executions."""
+
+ dag_bag: DBDagBag
+ session: Session
+ _dags: dict[UUID, SerializedDAG] = attrs.field(factory=dict, init=False)
+ _region_nodes: dict[UUID, str | None] = attrs.field(factory=dict,
init=False)
+
+ def get_task(
+ self, dag_id: str, run_id: str, task_id: str, *, dag_version_id: UUID
| None = None
+ ) -> SerializedOperator:
+ if dag_version_id is None:
+ return self._producer_task(dag_id, run_id, task_id)
+ if dag_version_id not in self._dags:
+ dag = self.dag_bag.get_dag(dag_version_id, session=self.session)
+ if dag is None or dag.dag_id != dag_id:
+ raise ValueError(f"Pinned Dag for {dag_id}/{run_id} not found")
+ self._dags[dag_version_id] = dag
+ return self._dags[dag_version_id].get_task(task_id)
+
+ def _producer_task(
+ self,
+ dag_id: str,
+ run_id: str,
+ task_id: str,
+ *,
+ region_id: UUID | None = None,
+ region_index: int | None = None,
+ ) -> SerializedOperator:
+ query = select(TaskInstance.dag_version_id).where(
+ TaskInstance.working_set.is_(True),
+ TaskInstance.dag_id == dag_id,
+ TaskInstance.run_id == run_id,
+ TaskInstance.task_id == task_id,
+ )
+ if region_id is not None:
+ query = query.where(TaskInstance.region_id == region_id)
+ if region_index is not None:
+ query = query.where(TaskInstance.region_index == region_index)
+ versions = set(self.session.scalars(query.distinct()))
+ if not versions or None in versions:
+ versions.discard(None)
+ version = self.session.scalar(
+ select(DagRun.created_dag_version_id).where(DagRun.dag_id ==
dag_id, DagRun.run_id == run_id)
+ )
+ if version is None:
+ raise ValueError(f"Pinned Dag for {dag_id}/{run_id} not found")
+ versions.add(version)
+ tasks = [self.get_task(dag_id, run_id, task_id,
dag_version_id=version) for version in versions]
+ classifications = {
+ (task.get_needs_expansion(), loop.group_id if (loop :=
enclosing_loop(task)) else None)
+ for task in tasks
+ }
+ if len(classifications) != 1:
+ raise AmbiguousProducerError("Producer definitions differ; select
explicit region coordinates")
+ return tasks[0]
+
+ def public_map_index(self, ti: TaskCoordinate, *, dag_version_id: UUID |
None = None) -> int:
+ if ti.region_id == SENTINEL_REGION_ID:
+ return ti.region_index
+ version = dag_version_id or getattr(ti, "dag_version_id", None)
+ try:
+ task = (
+ self.get_task(ti.dag_id, ti.run_id, ti.task_id,
dag_version_id=version)
+ if version is not None
+ else self._producer_task(
+ ti.dag_id, ti.run_id, ti.task_id, region_id=ti.region_id,
region_index=ti.region_index
+ )
+ )
+ except (TaskNotFound, ValueError):
+ return ti.region_index if self._is_mapped_region(ti) else -1
+ return ti.region_index if task.get_needs_expansion() else -1
+
+ def _is_mapped_region(self, ti: TaskCoordinate) -> bool:
+ """Tell from stored region data alone whether a task's own expansion
holds this coordinate."""
+ if ti.region_id not in self._region_nodes:
+ self._region_nodes[ti.region_id] = self.session.scalar(
+ select(DynamicRegion.node_id).where(DynamicRegion.id ==
ti.region_id)
+ )
+ return self._region_nodes[ti.region_id] == ti.task_id
+
+ def has_regions(self, dag_id: str, run_id: str | None, task_id: str) ->
bool:
+ query = select(TaskInstance.id).where(
+ TaskInstance.working_set.is_(True),
+ TaskInstance.dag_id == dag_id,
+ TaskInstance.task_id == task_id,
+ TaskInstance.region_id != SENTINEL_REGION_ID,
+ )
+ if run_id is not None:
+ query = query.where(TaskInstance.run_id == run_id)
+ return bool(self.session.scalar(select(query.exists())))
+
+ def resolve(
+ self,
+ *,
+ dag_id: str,
+ run_id: str,
+ task_id: str,
+ caller: TaskInstance | None = None,
+ region_id: UUID | None = None,
+ region_index: int | None = None,
+ map_indexes: int | Collection[int] | None = None,
+ previous_iteration: bool = False,
+ ) -> tuple[TaskInstance, ...]:
+ if region_index is not None and region_id is None:
+ raise ValueError("region_index requires an explicit producer
region_id")
+ if not previous_iteration and (
+ region_id == UUID(int=0) or (region_id is None and not
self.has_regions(dag_id, run_id, task_id))
Review Comment:
Nit: `UUID(int=0)` here and on the select below, while the rest of the
module uses the imported `SENTINEL_REGION_ID`.
##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py:
##########
@@ -545,14 +687,22 @@ def delete_xcom(
def _find_writer_id(
- attempt_id: UUID, *, dag_id: str, run_id: str, task_id: str, map_index:
int, session: SessionDep
+ attempt_id: UUID,
+ *,
+ dag_id: str,
+ run_id: str,
+ task_id: str,
+ map_index: int,
+ region_id: UUID | None = None,
+ region_index: int | None = None,
+ session: SessionDep,
) -> UUID | None:
"""Resolve the attempt a write targets: the caller's own, or the live
attempt of another task."""
coordinates = (
TaskInstance.dag_id == dag_id,
TaskInstance.run_id == run_id,
TaskInstance.task_id == task_id,
- TaskInstance.map_index == map_index,
+ *_coordinate_filters(map_index=map_index, region_id=region_id,
region_index=region_index),
Review Comment:
When another task calls `XCom.set(..., task_id="body")` or `XCom.delete` for
a loop member, every live pass matches
`public_map_index_expression(TaskInstance) == -1`, so the fallback
`session.scalar(select(TaskInstance.id).where(*coordinates))` picks whichever
pass the database returns first and writes or deletes there with no error. The
read path for the same coordinates raises `AmbiguousProducerError` and returns
409. Should the fallback fetch up to two ids and return 409 when more than one
matches, or go through `TaskCoordinateResolver.resolve` with the caller?
##########
airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_xcoms.py:
##########
@@ -46,6 +52,170 @@
pytestmark = pytest.mark.db_test
[email protected]("suffix", ["", "/item/0", "/item/-1", "/slice"])
+def
test_regional_mapped_xcom_reads_resolve_live_producers_before_slicing(client,
dag_maker, session, suffix):
+ with dag_maker(serialized=True):
+ PythonOperator.partial(task_id="mapped",
python_callable=str).expand(op_args=[[1], [2]])
+ dr = dag_maker.create_dagrun()
+ first = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id, node_id="mapped")
+ session.add(first)
+ session.flush()
+ replacement = DynamicRegion(
+ dag_id=dr.dag_id, run_id=dr.run_id, node_id="mapped",
forked_from_region_id=first.id
+ )
+ session.add(replacement)
+ session.flush()
+ tis = sorted(dr.task_instances, key=lambda ti: ti.region_index)
+ tis[0].region_id = first.id
+ tis[1].region_id = replacement.id
+ superseded = TaskInstance(
+ tis[0].task, tis[0].dag_version_id, run_id=dr.run_id, map_index=1,
region_id=first.id
+ )
+ session.add(superseded)
+ session.flush()
+ for ti, value in [(tis[0], "zero"), (superseded, "archived"), (tis[1],
"one")]:
+ XComModel.set_for_attempt(
+ task_instance_id=ti.id, key="key", value=value, serialize=False,
session=session
+ )
+ superseded.archive(reason="test", session=session)
+ session.commit()
+ url = f"/execution/xcoms/{dr.dag_id}/{dr.run_id}/mapped/key"
+
+ assert client.head(url).headers["Content-Range"] == "map_indexes 2"
+ response = client.get(url + suffix, params={"map_index": 0} if not suffix
else {})
+ assert response.status_code == 200
+ expected = {
+ "": {"key": "key", "value": "zero"},
+ "/item/0": "zero",
+ "/item/-1": "one",
+ "/slice": ["zero", "one"],
+ }
+ assert response.json() == expected[suffix]
+
+
+def test_xcom_write_coordinates_must_match_the_calling_task_instance(
+ client, create_task_instance, authenticate_as, session
+):
+ ti = create_task_instance()
+ region = DynamicRegion(dag_id=ti.dag_id, run_id=ti.run_id, node_id="loop")
+ session.add(region)
+ session.commit()
+ authenticate_as(ti)
+ url = f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/key"
+ other = {"region_id": str(region.id), "region_index": -1}
+ own = {"region_id": str(ti.region_id), "region_index": ti.region_index}
+
+ assert client.post(url, params=other, json="regional").status_code == 404
+ assert client.delete(url, params=other).status_code == 404
+ assert client.post(url, params=own, json="legacy").status_code == 201
+ assert client.get(url, params=own).json() == {"key": "key", "value":
"legacy"}
+ assert client.get(url, params=other).status_code == 404
+ assert client.delete(url, params=own).status_code == 200
+
+
[email protected]
+def loop_xcoms(client, dag_maker, session):
+ @task_group
+ def body():
+ EmptyOperator(task_id="producer") >> EmptyOperator(task_id="consumer")
+
+ with dag_maker(serialized=True) as dag:
+ EmptyOperator(task_id="outside") >> create_loop(body, max_iterations=3)
+ dr = dag_maker.create_dagrun()
+ first = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id, node_id="body")
+ session.add(first)
+ session.flush()
+ replacement = DynamicRegion(
+ dag_id=dr.dag_id, run_id=dr.run_id, node_id="body",
forked_from_region_id=first.id
+ )
+ session.add(replacement)
+ session.flush()
+ tis = {ti.task_id: ti for ti in dr.task_instances}
+ producer, consumer = tis["body.producer"], tis["body.consumer"]
+ producer.region_id, producer.region_index = first.id, 2
+ consumer.region_id, consumer.region_index = replacement.id, 2
+ previous = TaskInstance(
+ task=dag.get_task(producer.task_id), run_id=dr.run_id,
dag_version_id=producer.dag_version_id
+ )
+ previous.region_id, previous.region_index = first.id, 1
+ session.add(previous)
+ session.flush()
+ for ti, value in [(producer, "current"), (previous, "previous"),
(tis["outside"], "outside")]:
+ XComModel.set_for_attempt(
+ task_instance_id=ti.id, key="key", value=value, serialize=False,
session=session
+ )
+ session.commit()
+ exec_app = client.app.routes[-1].app
Review Comment:
The `authenticate_as` fixture in this file already does this override (and
the test just above uses it), and the `client` fixture pops `require_auth` at
teardown. Could `loop_xcoms` take `authenticate_as`, call
`authenticate_as(consumer)` and drop the `routes[-1]` lookup and the manual
restore?
##########
airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py:
##########
@@ -1391,45 +1457,72 @@ def get_task_instance_states(
if run_ids:
query = query.where(TI.run_id.in_(run_ids))
- if map_index is not None:
- query = query.where(TI.map_index == map_index)
-
- results = session.scalars(query).all()
+ query = _filter_task_coordinates(
+ query, map_index=map_index, region_id=region_id,
region_index=region_index
+ )
+ results = session.execute(query).all()
if task_group_id:
group_tasks = _get_group_tasks(
dag_id, task_group_id, session, dag_bag, logical_dates, run_ids,
map_index
)
- results = results + group_tasks if task_ids else group_tasks
-
- [
- run_id_task_state_map[task.run_id].update(
- {task.task_id: task.state}
- if task.map_index < 0
- else {f"{task.task_id}_{task.map_index}": task.state}
+ group_query = _filter_task_coordinates(
+ select(TI, public_map_index_expression(TI)).where(TI.id.in_(ti.id
for ti in group_tasks)),
+ map_index=map_index,
+ region_id=region_id,
+ region_index=region_index,
)
- for task in results
- ]
+ group_results = session.execute(group_query).all()
+ results = [*results, *group_results] if task_ids else group_results
+
+ identities: dict[tuple[str, str], UUID] = {}
+ for task, public_index in results:
+ key = task.task_id if public_index < 0 else
f"{task.task_id}_{public_index}"
+ identity = task.run_id, key
+ if identity in identities and identities[identity] != task.id:
+ raise HTTPException(status.HTTP_409_CONFLICT, "Task states require
explicit region coordinates")
+ identities[identity] = task.id
+ run_id_task_state_map[task.run_id][key] = task.state
return TaskStatesResponse(task_states=run_id_task_state_map)
@router.get("/breadcrumbs", status_code=status.HTTP_200_OK)
async def get_task_instance_breadcrumbs(
- dag_id: str, run_id: str, session: AsyncSessionDep
+ dag_id: str,
+ run_id: str,
+ session: AsyncSessionDep,
+ region_id: UUID | None = None,
+ region_index: int | None = None,
) -> TaskBreadcrumbsResponse:
+ query = (
+ select(
+ TI.task_id,
+ public_map_index_expression(TI).label("map_index"),
+ TI.state,
+ TI.operator,
+ TI.duration,
+ TI.region_id,
+ TI.region_index,
+ )
+ .where(TI.working_set.is_(True))
+ .where(TI.dag_id == dag_id, TI.run_id == run_id,
TI.state.in_(TerminalTIState))
+ .order_by(TI.task_id, public_map_index_expression(TI), TI.region_id,
TI.region_index)
+ )
result = (
await session.execute(
- select(TI.task_id, TI.map_index, TI.state, TI.operator,
TI.duration)
- .where(TI.dag_id == dag_id, TI.run_id == run_id,
TI.state.in_(TerminalTIState))
- .order_by(TI.task_id, TI.map_index)
+ _filter_task_coordinates(query, map_index=None,
region_id=region_id, region_index=region_index)
)
).mappings()
def _iter_breadcrumbs() -> Iterator[dict[str, Any]]:
for row in result:
- yield {str(k): v for k, v in row.items()}
+ yield {
+ str(k): v
+ for k, v in row.items()
+ if row.region_id != SENTINEL_REGION_ID or k not in
{"region_id", "region_index"}
Review Comment:
Breadcrumbs drop `region_id`/`region_index` for a non-regional TI, but
`/previous` returns `region_id=ti.region_id`, which is the zero UUID for every
ordinary task, and the model default is `None`. A client reading the spec can't
tell which of the three means "not in a region" without knowing the internal
sentinel. Would it be worth mapping the sentinel to `None` at the API boundary
for all these responses before the version ships?
##########
task-sdk/src/airflow/sdk/execution_time/schema/versions/v2026_10_30.py:
##########
@@ -29,8 +29,60 @@
from cadwyn import VersionChange, schema
from airflow.dag_processing.processor import DagFileParsingResult # noqa:
SDK002
-from airflow.sdk.api.datamodels._generated import TIRunContext
-from airflow.sdk.execution_time.comms import TaskState
+from airflow.sdk.api.datamodels._generated import PreviousTIResponse,
TaskInstance, TIRunContext
+from airflow.sdk.execution_time.comms import (
+ DeleteXCom,
+ GetPreviousTI,
+ GetTaskBreadcrumbs,
+ GetTaskStates,
+ GetTICount,
+ GetXCom,
+ GetXComCount,
+ GetXComSequenceItem,
+ GetXComSequenceSlice,
+ SetXCom,
+ TaskState,
+)
+
+
+class AddRegionSelectors(VersionChange):
+ """Carry explicit coordinates in XCom and task instance requests."""
+
+ description = __doc__
+ instructions_to_migrate_to_previous_version = (
+ schema(GetTICount).field("region_id").didnt_exist,
+ schema(GetTICount).field("region_index").didnt_exist,
+ schema(GetTaskStates).field("region_id").didnt_exist,
+ schema(GetTaskStates).field("region_index").didnt_exist,
+ schema(GetPreviousTI).field("region_id").didnt_exist,
+ schema(GetPreviousTI).field("region_index").didnt_exist,
+ schema(GetTaskBreadcrumbs).field("region_id").didnt_exist,
+ schema(GetTaskBreadcrumbs).field("region_index").didnt_exist,
+ schema(GetXCom).field("region_id").didnt_exist,
+ schema(GetXCom).field("region_index").didnt_exist,
+ schema(GetXComCount).field("region_id").didnt_exist,
+ schema(GetXComCount).field("region_index").didnt_exist,
+ schema(GetXComSequenceItem).field("region_id").didnt_exist,
+ schema(GetXComSequenceItem).field("region_index").didnt_exist,
+ schema(GetXComSequenceSlice).field("region_id").didnt_exist,
+ schema(GetXComSequenceSlice).field("region_index").didnt_exist,
+ schema(SetXCom).field("region_id").didnt_exist,
+ schema(SetXCom).field("region_index").didnt_exist,
+ schema(DeleteXCom).field("region_id").didnt_exist,
+ schema(DeleteXCom).field("region_index").didnt_exist,
+ )
+
+
+class AddTaskInstanceRegionCoordinates(VersionChange):
Review Comment:
This shares its name with the Execution API change
`AddTaskInstanceRegionCoordinates` in
`airflow-core/.../versions/v2026_10_30.py`, though the docstring on
`AddArgBindingsToSupervisorTIRunContext` below says supervisor mirrors are
named apart so the two migrations aren't confused. Maybe
`AddRegionCoordinatesToSupervisorTaskInstance`?
--
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]