kaxil commented on code in PR #74346:
URL: https://github.com/apache/airflow/pull/74346#discussion_r4198201258
##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -2092,38 +2135,50 @@ def create_ti(task: Operator, indexes: Iterable[int])
-> Iterator[TI]:
def _create_tasks(
self,
tasks: Iterable[Operator],
- task_creator: Callable[[Operator, Iterable[int]], CreatedTasks],
+ task_creator: Callable[[Operator, Iterable[int], UUID], CreatedTasks],
*,
session: Session,
+ parent_region: tuple[UUID, int] | None = None,
+ expand_literals: bool = False,
) -> CreatedTasks:
"""
- Create missing tasks -- and expand any MappedOperator that _only_ have
literals as input.
+ Create ordinary tasks and a region-owned placeholder for each mapped
task.
:param tasks: Tasks to create jobs for in the DAG run
:param task_creator: Function to create task instances
+ :param expand_literals: Create the slots of a task whose inputs are
all literals right away,
+ so its placeholder is never expanded one task at a time.
"""
- from airflow.models.expandinput import NotFullyPopulated
- from airflow.serialization.definitions.mappedoperator import
get_mapped_ti_count
-
- map_indexes: Iterable[int]
+ tasks = list(tasks)
+ regions = {
+ task.task_id: DynamicRegion(
+ dag_id=self.dag_id,
+ run_id=self.run_id,
+ node_id=task.task_id,
+ parent_region_id=parent_region[0] if parent_region else None,
+ parent_region_index=parent_region[1] if parent_region else
None,
+ )
+ for task in tasks
+ if task.get_needs_expansion()
+ }
+ session.add_all(regions.values())
+ session.flush()
for task in tasks:
- try:
- count = get_mapped_ti_count(task, self.run_id, session=session)
- except (NotMapped, NotFullyPopulated):
- map_indexes = (-1,)
+ indexes: Iterable[int]
+ if task.task_id in regions:
+ region_id, indexes = regions[task.task_id].id, (-1,)
Review Comment:
Once the placeholder is born in the task's own region, nothing moves it back
if a later Dag version stops mapping the task.
`_check_for_removed_or_restored_tasks` still does `except NotMapped: pass`, so
on an unversioned bundle the repinned rows stay at `(R, -1)` (or `(R, 0..n)`)
under a definition that calls the task plain. Two things then break that worked
while these rows sat in the sentinel: `prepare_task_log_contexts` builds
`node_kinds` from the pinned Dag, which no longer lists the task as `"map"`, so
`region_log_position` raises `KeyError` inside the scheduler's enqueue critical
section (and the same TI is picked again after a restart); and
`resolve_current_producers` skips non-sentinel rows when `is_mapped` is false,
so a downstream XCom pull or `.expand()` over this task finds no producer.
Could the `NotMapped` branch move an unfinished own-region placeholder back to
`SENTINEL_REGION_ID` and mark its `>= 0` slots removed? Classifying an own
top-level region from stored d
ata in the log path, the way `_is_mapped_region` already does, would also keep
one stale row from taking the scheduler down.
##########
airflow-core/src/airflow/serialization/definitions/xcom_arg.py:
##########
@@ -174,58 +173,32 @@ def _(
task_id = xcom_arg.operator.task_id
mapped = is_mapped(xcom_arg.operator)
- if producer_contexts is not None:
- producers = resolve_current_producers(
- dag_id=dag_id,
- run_id=run_id,
- task_id=task_id,
- is_mapped=mapped,
- context=producer_contexts.get(task_id),
- session=session,
- )
- if not producers:
- return None
- if mapped and any(ti.state in State.unfinished for ti in producers):
- return None
- read = XComModel.get_many(
- dag_ids=dag_id,
- run_id=run_id,
- task_ids=task_id,
- key=XCOM_RETURN_KEY,
-
producer_ids=select(TaskInstance.id).where(TaskInstance.id.in_([ti.id for ti in
producers])),
- )
- elif mapped:
- unfinished_ti_exists = exists_query(
- TaskInstance.working_set.is_(True),
- TaskInstance.dag_id == dag_id,
- TaskInstance.run_id == run_id,
- TaskInstance.task_id == task_id,
- TaskInstance.region_id == SENTINEL_REGION_ID,
- # Special NULL treatment is needed because 'state' can be NULL.
- # The "IN" part would produce "NULL NOT IN ..." and eventually
- # "NULl = NULL", which is a big no-no in SQL.
- or_(
- TaskInstance.state.is_(None),
- TaskInstance.state.in_(s.value for s in State.unfinished if s
is not None),
- ),
- session=session,
- )
- if unfinished_ti_exists:
- return None # Not all of the expanded tis are done yet.
- read = XComModel.get_many(dag_ids=dag_id, run_id=run_id,
task_ids=task_id, key=XCOM_RETURN_KEY)
- else:
- read = XComModel.get_many(
- dag_ids=dag_id, run_id=run_id, task_ids=task_id, map_indexes=-1,
key=XCOM_RETURN_KEY
- )
+ producers = resolve_current_producers(
Review Comment:
This drops the old no-context branch, so every call now loads every live
producer TI as a full ORM row, looks up its region, and then sends an `IN` list
with one UUID per producer. For `tg.expand(x=mapped.output)` that runs from
`_get_expanded_ti_count` once per pending group instance per pass, so an N-wide
producer feeding an M-wide group loads about N x M rows per scheduling pass,
where before it ran an EXISTS and a count per instance. Could the path without
a loop context keep the SQL-side EXISTS and count, with `public_region_filter`
in place of the sentinel filter? Failing that, selecting only `id` and `state`
would at least avoid the hydration.
##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -2092,38 +2135,50 @@ def create_ti(task: Operator, indexes: Iterable[int])
-> Iterator[TI]:
def _create_tasks(
self,
tasks: Iterable[Operator],
- task_creator: Callable[[Operator, Iterable[int]], CreatedTasks],
+ task_creator: Callable[[Operator, Iterable[int], UUID], CreatedTasks],
*,
session: Session,
+ parent_region: tuple[UUID, int] | None = None,
+ expand_literals: bool = False,
) -> CreatedTasks:
"""
- Create missing tasks -- and expand any MappedOperator that _only_ have
literals as input.
+ Create ordinary tasks and a region-owned placeholder for each mapped
task.
:param tasks: Tasks to create jobs for in the DAG run
:param task_creator: Function to create task instances
+ :param expand_literals: Create the slots of a task whose inputs are
all literals right away,
+ so its placeholder is never expanded one task at a time.
"""
- from airflow.models.expandinput import NotFullyPopulated
- from airflow.serialization.definitions.mappedoperator import
get_mapped_ti_count
-
- map_indexes: Iterable[int]
+ tasks = list(tasks)
+ regions = {
Review Comment:
Before this change, two transactions that both created the placeholder for a
newly added mapped task (or both expanded the same sentinel placeholder)
collided on `task_instance_current_key` and one rolled back in
`_create_task_instances`. Now each mints its own `DynamicRegion` with a fresh
id, and `dynamic_region` has no unique key on the top-level slot, so both
commits succeed. Paths that can overlap without holding the DagRun row lock are
`clear(only_new=True)` through `_update_dagrun_to_latest_version` racing the
scheduler, and `_update_dag_run_state_for_paused_dags` on two HA schedulers.
The outcome is two live expansions of one task, after which public lookups
raise `MultipleResultsFound` or `AmbiguousProducerError`. Should these paths
take the DagRun row lock, or should top-level unforked regions get a uniqueness
guard per `node_id`?
##########
airflow-core/src/airflow/utils/log/task_log_address.py:
##########
@@ -175,15 +175,22 @@ def prepare_task_log_contexts(
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
+ version_id
+ for ti in regional
+ if (version_id := ti.dag_version_id or runs[ti.dag_id,
ti.run_id].created_dag_version_id)
}
- 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] = {
+ attached = {
+ dag.dag_version_id: dag
+ for ti in regional
+ if (dag := getattr(ti.task, "dag", None)) is not None and
dag.dag_version_id is not None
+ }
+ dag_bag = DBDagBag(load_op_links=False)
Review Comment:
With every mapped expansion in its own region, every mapped TI now takes
this branch, and a fresh `DBDagBag` per call means each scheduler enqueue batch
deserializes each pinned Dag version again, inside the pool-locked critical
section (scheduler TIs carry no `ti.task`, so `attached` is empty). Before
this, only loop TIs paid for it. For a task's own top-level region the
`DynamicRegion` rows loaded just above already say it is a map with no loop
position, so the Dag is only needed when an ancestor region belongs to a loop.
`get_shared_dag_bag()` from this PR looks meant for exactly this, but nothing
calls it yet; either wire it in here or drop it.
##########
airflow-core/src/airflow/ti_deps/dep_context.py:
##########
@@ -106,6 +109,73 @@ class DepContext:
every ``UP_FOR_RESCHEDULE`` task instance. With ``init=False`` those
instances would each get a
fresh empty dict, so they would neither read the memo nor warm it for
anything else.
"""
+ regional_runs: dict[tuple[str, str], bool] = attr.ib(factory=dict,
repr=False)
+ producer_tis: dict[tuple[str, str, UUID | None, UUID, int, str],
tuple[TaskInstance, ...]] = attr.ib(
+ factory=dict, repr=False
+ )
+ coordinate_resolvers: dict[Session, TaskCoordinateResolver] =
attr.ib(factory=dict, repr=False)
+ dynamic_dags: dict[int, bool] = attr.ib(factory=dict, repr=False)
+
+ def has_regions(self, ti: TaskInstance, *, session: Session) -> bool:
+ from airflow.models.dynamic_region import DynamicRegion
+
+ key = ti.dag_id, ti.run_id
+ if key not in self.regional_runs:
+ self.regional_runs[key] = self._dag_can_have_regions(ti) and bool(
+ session.scalar(
+ select(
+ select(DynamicRegion.id)
+ .where(
+ DynamicRegion.dag_id == ti.dag_id,
+ DynamicRegion.run_id == ti.run_id,
+ )
+ .exists()
+ )
+ )
+ )
+ return self.regional_runs[key]
+
+ def _dag_can_have_regions(self, ti: TaskInstance) -> bool:
+ from airflow.models.task_coordinates import enclosing_loop
+
+ dag = getattr(ti.task, "dag", None)
+ if dag is None:
+ return True
+ if id(dag) not in self.dynamic_dags:
+ self.dynamic_dags[id(dag)] = any(
+ task.get_needs_expansion() or enclosing_loop(task) is not None
Review Comment:
Since every mapped task now gets a region at run creation, this makes
`has_regions` true for every run of any Dag with a single `.expand()`, not only
loop Dags. In that mode `_is_relevant_upstream` calls `upstream_tis` for each
finished upstream, which costs an EXISTS plus a SELECT per distinct upstream
task (three queries for a mapped upstream) on every scheduling pass, and the
memo is cleared on each mid-pass expansion or revision. The legacy branch
answered a plain fan-in from `finished_tis` with no queries. Without loops each
task has one live expansion, so the legacy predicates should still be exact;
could this gate on `enclosing_loop(task) is not None` only, or skip resolution
for upstreams that are neither mapped nor in a loop?
##########
airflow-core/src/airflow/models/taskinstance.py:
##########
@@ -1463,9 +1486,12 @@ def field(name: str) -> Any:
return "<deferred>"
prefix = f"<TaskInstance: {field('dag_id')}.{field('task_id')}
{field('run_id')} "
- map_index = field("map_index")
- if map_index != -1:
- prefix += f"map_index={map_index} "
+ region_id = field("region_id")
+ region_index = field("region_index")
+ if region_id != SENTINEL_REGION_ID:
Review Comment:
Every new mapped expansion now lives outside the sentinel, so for ordinary
mapped tasks with no loop this branch replaces `map_index=2` with
`region_id=<uuid> region_index=2`. The same switch is in `_log_state` ("Marking
task as ..."), the scheduler's "TaskInstance Finished" line,
`TaskInstanceNote.__repr__` and the heartbeat-timeout message details, where
`Map Index` becomes `Region Id` / `Region Index` in what users see in failure
callbacks. Could these branch on whether the coordinate is the task's own map
region and keep the `map_index` form there?
##########
airflow-core/src/airflow/models/dagrun.py:
##########
@@ -1981,18 +2012,23 @@ def _check_for_removed_or_restored_tasks(
except NotFullyPopulated:
# What if it is _now_ dynamically mapped, but wasn't before?
try:
- total_length = get_mapped_ti_count(task, self.run_id,
session=session)
+ total_length = get_mapped_ti_count(
+ task,
+ self.run_id,
+ session=session,
+ producer_contexts=coordinates.producer_contexts(ti),
Review Comment:
`producer_contexts(ti)` is only non-empty for loop producers, but here it
runs for every TI of an XCom-mapped task, and these TIs come from
`get_task_instances` without `ti.task`, so it resolves the definition through
`ti.dag_version_id` and then `created_dag_version_id`. For a run carried over
from Airflow 2 both are NULL on finished TIs, `_producer_task` raises
`ValueError`, and `verify_integrity` aborts. The scheduler logs that and still
commits the version bump on the unfinished TIs, so integrity is never
re-verified for that version. The loop already has `task` from `dag`; could you
pass that in (or set `ti.task = task` first) so this needs no pinned-version
lookup?
##########
ts-sdk/src/generated/dag-schema-fields.ts:
##########
@@ -113,6 +113,8 @@ export interface GeneratedTaskFields {
readonly executor?: string;
/** Maps to the schema key `do_xcom_push` (schema default `true`). */
readonly doXcomPush?: boolean;
+ /** Maps to the schema key `returns_dag_result` (schema default `false`). */
+ readonly returnsDagResult?: boolean;
Review Comment:
Adding `returns_dag_result` to the operator schema turns it into an
authoring option in the TypeScript SDK here, `TaskSpec.ReturnsDagResult` in Go
(`spec.gen.go`), and, as far as I can tell, a generated attribute in the Java
DSL too. None of those runtimes sends `dag_result` on the return-value XCom
push, so a TS author who sets `returnsDagResult: true` gets a run whose wait
endpoint returns no results and no error. Both generators say an attribute the
SDK can't honour should be excluded rather than shipped, since it is hard to
take away later. Could you add `returns_dag_result` to `EXCLUDED_TASK_FIELDS`,
the Go `taskShape` exclude map and `excludedTaskKeys`, then regenerate? Core
keeps the schema property either way.
##########
airflow-core/src/airflow/cli/commands/task_command.py:
##########
@@ -389,7 +447,19 @@ def task_states_for_dag_run(args, *, session: Session =
NEW_SESSION) -> None:
"not found"
)
- has_mapped_instances = any(ti.map_index >= 0 for ti in
dag_run.task_instances)
+ rows = session.execute(
+ select(TaskInstance, public_map_index_expression(TaskInstance)).where(
+ TaskInstance.working_set.is_(True),
+ TaskInstance.dag_id == dag_run.dag_id,
+ TaskInstance.run_id == dag_run.run_id,
+ )
+ ).all()
+ task_instances = [ti for ti, _ in rows]
+ map_indexes = {ti.id: map_index for ti, map_index in rows}
+ has_mapped_instances = any(map_index >= 0 for map_index in
map_indexes.values())
+ has_loop_instances = any(
+ ti.region_id != SENTINEL_REGION_ID and map_indexes[ti.id] < 0 for ti
in task_instances
Review Comment:
An unexpanded XCom-mapped placeholder (or a zero-length expansion's skipped
placeholder) now sits at `region_index = -1` in its own region, and
`public_map_index_expression` returns -1 for it, so this flags a plain mapped
Dag as having loop instances and adds `region_id` / `region_index` to every row
of the output. The new test only passes because it uses a literal expansion,
which is created already expanded. Could this check that the row's region is
not the task's own top-level region instead of relying on `map_index < 0`, and
cover the unexpanded XCom-mapped case in the test?
##########
airflow-core/src/airflow/jobs/scheduler_job_runner.py:
##########
@@ -3971,7 +3973,15 @@ def _purge_task_instances_without_heartbeats(
has_callback_version = _ensure_ti_has_dag_version_id(ti, session,
self.log)
if not has_callback_version and ti.state !=
TaskInstanceState.RESTARTING:
continue
+ callback_map_index = None
if has_callback_version:
+ try:
+ callback_map_index = coordinates.public_map_index(ti)
+ except (TaskNotFound, ValueError):
Review Comment:
`public_map_index` already catches `TaskNotFound` and `ValueError`
(including `AmbiguousProducerError`) and falls back to `_is_mapped_region`, so
this `except` can't fire and `callback_map_index` is never `None` once
`has_callback_version` is true. When the definition really is missing, the
callback still goes out with a guessed index, the opposite of what the warning
says. Either drop the try/except, or check `coordinates.get_task(...)`
explicitly if skipping is the intent (and assert that in
`test_heartbeat_timeout_completes_clear_without_definition`).
##########
airflow-core/src/airflow/models/taskinstance.py:
##########
@@ -2587,7 +2631,14 @@ def expand_mapped_task(self, *, session: Session) ->
tuple[Sequence[TaskInstance
)
try:
- total_length: int | None = get_mapped_ti_count(task, run_id,
session=session)
+ total_length: int | None = get_mapped_ti_count(
+ task,
+ run_id,
+ session=session,
+ producer_contexts=TaskCoordinateResolver.for_dag(
+ getattr(task, "dag", None), session
Review Comment:
`task` is already known to be a `SerializedMappedOperator` or
`SerializedBaseOperator` here, and the `except` block below reads `task.dag`
directly, so the `getattr` isn't needed.
```suggestion
task.dag, session
```
##########
airflow-core/tests/unit/models/test_mappedoperator.py:
##########
@@ -57,6 +63,97 @@
from airflow.sdk.definitions.context import Context
[email protected]("mapping", ["dict", "list", "group"])
+def test_mapped_count_uses_retained_producer_in_callers_loop_pass(dag_maker,
session, mapping):
+ @task_group
+ def body():
+ source = PythonOperator(task_id="source", python_callable=list)
+ if mapping == "dict":
+ PythonOperator.partial(task_id="consumer",
python_callable=list).expand(op_kwargs=source.output)
+ elif mapping == "list":
+ PythonOperator.partial(task_id="consumer",
python_callable=list).expand_kwargs(source.output)
+ else:
+
+ @task_group
+ def mapped_group(value):
+ PythonOperator(task_id="consumer", python_callable=list)
+
+ mapped_group.expand(value=source.output)
+
+ with dag_maker(serialized=True):
+ loop = create_loop(body, max_iterations=3)
+ dr = dag_maker.create_dagrun()
+ original = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id,
node_id=loop.group_id)
+ session.add(original)
+ session.flush()
+ fork = DynamicRegion(
+ dag_id=dr.dag_id,
+ run_id=dr.run_id,
+ node_id=loop.group_id,
+ forked_from_region_id=original.id,
+ resumes_from_index=1,
+ )
+ session.add(fork)
+ session.flush()
+ source = next(ti for ti in dr.task_instances if ti.task_id ==
"body.source")
+ consumer = next(ti for ti in dr.task_instances if
ti.task_id.endswith("consumer"))
+ source.region_id, source.region_index = original.id, 1
+ source.state = TaskInstanceState.SUCCESS
+ expansion = DynamicRegion(
+ dag_id=dr.dag_id,
+ run_id=dr.run_id,
+ node_id=consumer.task_id,
+ parent_region_id=fork.id,
+ parent_region_index=1,
+ )
+ session.add(expansion)
+ session.flush()
+ consumer.region_id, consumer.region_index = expansion.id, 0
+ earlier = TaskInstance(
+ dag_maker.serialized_dag.get_task(source.task_id),
+ dag_version_id=source.dag_version_id,
+ run_id=dr.run_id,
+ region_id=original.id,
+ region_index=0,
+ state=TaskInstanceState.SUCCESS,
+ )
+ session.add(earlier)
+ session.flush()
+ for producer, length in [(earlier, 99), (source, 2)]:
+ XComModel.set_for_attempt(
+ task_instance_id=producer.id,
+ key="return_value",
+ value=[{}] * length,
+ serialize=False,
+ mapped_length=length,
+ session=session,
+ )
+ session.flush()
+ resolver = TaskCoordinateResolver(DBDagBag(), session)
+ contexts = resolver.producer_contexts(consumer)
+
+ assert (
+ get_mapped_ti_count(
+ dag_maker.serialized_dag.get_task(consumer.task_id),
+ dr.run_id,
+ producer_contexts=contexts,
+ session=session,
+ )
+ == 2
+ )
+ if mapping == "group":
+ consumer.task = dag_maker.serialized_dag.get_task(consumer.task_id)
+ assert (
+ consumer.get_relevant_upstream_map_indexes(
+ consumer.task,
+ 2,
+ producer_contexts=contexts,
+ session=session,
+ )
+ == 0
Review Comment:
With the consumer at `region_index` 0, `_get_relevant_map_indexes` returns
`0 * ancestor_ti_count // 2`, which is 0 whether the group count comes from the
retained pass-1 producer (2) or the earlier pass-0 one (99). So this half can't
catch a wrong producer. Setting `consumer.region_index = 1` and asserting `==
1` would (the wrong producer gives 49).
##########
airflow-core/tests/unit/ti_deps/deps/test_trigger_rule_dep.py:
##########
@@ -2226,9 +2401,18 @@ def plain():
session.commit()
return dr
- def test_memoized_across_downstreams_sharing_upstream(self, dag_maker,
session):
+ @pytest.mark.parametrize("regional", [False, True])
Review Comment:
`src` is a literal expansion, so `create_dagrun()` already puts it in its
own region and `regional=False` runs the regional path too; the `True` branch
only adds a second top-level region for the same node, which production never
creates. That leaves nothing in this class covering a pre-region (sentinel)
expansion, which the commit says stays in production. Could `regional=False`
move the `src` TIs to `SENTINEL_REGION_ID` and delete the run's regions, as
`test_trigger_count_cache_separates_expansions_in_different_loop_passes` does,
and the `True` branch drop the hand-made region?
`test_mapped_siblings_share_resolved_producers` has the same duplicate-region
setup.
##########
airflow-core/tests/unit/cli/commands/test_task_command.py:
##########
@@ -74,6 +77,243 @@ def reset(dag_id):
session.execute(delete(SerializedDagModel).where(SerializedDagModel.dag_id ==
dag_id))
+def
test_task_states_for_dag_run_projects_map_index_and_identifies_regions(dag_maker,
session):
+ with dag_maker() as dag:
+ BashOperator(task_id="work", bash_command="echo work")
+ BashOperator.partial(task_id="mapped").expand(bash_command=["echo
mapped"])
+ run = dag_maker.create_dagrun()
+ loop = DynamicRegion(dag_id=run.dag_id, run_id=run.run_id, node_id="loop")
+ successor = DynamicRegion(dag_id=run.dag_id, run_id=run.run_id,
node_id="loop")
+ mapped = DynamicRegion(dag_id=run.dag_id, run_id=run.run_id,
node_id="mapped")
+ session.add_all([loop, successor, mapped])
+ session.flush()
+ work_ti = next(ti for ti in run.task_instances if ti.task_id == "work")
+ mapped_ti = next(ti for ti in run.task_instances if ti.task_id == "mapped")
+ work_ti.region_id, work_ti.region_index = loop.id, 4
+ mapped_ti.region_id, mapped_ti.region_index = mapped.id, 0
+ sibling = TaskInstance(
+ dag.get_task("work"),
+ run_id=run.run_id,
+ dag_version_id=work_ti.dag_version_id,
+ region_id=successor.id,
+ region_index=4,
+ )
+ session.add(sibling)
+ session.flush()
+
+ with redirect_stdout(io.StringIO()) as stdout:
+ task_command.task_states_for_dag_run(
+ cli_parser.get_parser().parse_args(
+ ["tasks", "states-for-dag-run", run.dag_id, run.run_id,
"--output", "json"]
+ ),
+ session=session,
+ )
+
+ rows = json.loads(stdout.getvalue())
+ assert len(rows) == 3
+ assert {(row["region_id"], row["region_index"], row["map_index"]) for row
in rows} == {
+ (str(loop.id), "4", ""),
+ (str(successor.id), "4", ""),
+ (str(mapped.id), "0", "0"),
+ }
+
+
+def
test_task_states_for_dag_run_lists_only_live_rows_without_region_columns_for_mapped_dag(
+ dag_maker, session
+):
+ with dag_maker():
+ BashOperator(task_id="work", bash_command="echo work")
+ BashOperator.partial(task_id="mapped").expand(bash_command=["echo a",
"echo b"])
+ run = dag_maker.create_dagrun()
+ work_ti = next(ti for ti in run.task_instances if ti.task_id == "work")
+ work_ti.state = State.SUCCESS
+ work_ti.archive(reason="cleared", session=session)
+ session.flush()
+ session.expire(run, ["task_instances"])
+
+ with redirect_stdout(io.StringIO()) as stdout:
+ task_command.task_states_for_dag_run(
+ cli_parser.get_parser().parse_args(
+ ["tasks", "states-for-dag-run", run.dag_id, run.run_id,
"--output", "json"]
+ ),
+ session=session,
+ )
+
+ rows = json.loads(stdout.getvalue())
+ assert sorted((row["task_id"], row["map_index"]) for row in rows) ==
[("mapped", "0"), ("mapped", "1")]
+ assert all(
+ set(row) == {"dag_id", "logical_date", "task_id", "state",
"start_date", "end_date", "map_index"}
+ for row in rows
+ )
+
+
[email protected]("command", ["state", "failed-deps", "test", "render"])
+def test_cli_commands_reuse_existing_regional_mapped_task(dag_maker, session,
mocker, capsys, command):
+ with dag_maker(dag_id="regional_cli", serialized=True) as dag:
+ PythonOperator.partial(task_id="mapped",
python_callable=str).expand(op_args=[[1], [2]])
+ dr = dag_maker.create_dagrun(run_id="regional")
+ region = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id,
node_id="mapped")
+ session.add(region)
+ session.flush()
+ for ti in dr.task_instances:
+ ti.region_id = region.id
+ ti.state = State.SUCCESS
+ selected = next(ti for ti in dr.task_instances if ti.region_index == 0)
+ selected_id = selected.id
+ session.commit()
+ serialized = dag_maker.serialized_dag
+ mocker.patch.object(task_command, "get_db_dag", autospec=True,
return_value=serialized)
+ mocker.patch.object(task_command, "get_bagged_dag", autospec=True,
return_value=dag_maker.dag)
+ run_task = mocker.patch.object(task_command, "_run_task", autospec=True)
+ lookup = mocker.spy(task_command, "_get_ti")
+ args = cli_parser.get_parser().parse_args(
+ ["tasks", command, dag.dag_id, "mapped", dr.run_id, "--map-index", "0"]
+ )
+
+ getattr(task_command, f"task_{command.replace('-', '_')}")(args)
+
+ assert lookup.spy_return[0].id == selected_id
+ if command == "test":
+ assert run_task.call_args.kwargs["ti"].id == selected_id
+ elif command == "state":
+ assert "success" in capsys.readouterr().out
+ session.expire_all()
+ rows = session.scalars(select(TaskInstance).where(TaskInstance.dag_id ==
dag.dag_id)).all()
+ assert len(rows) == 2
+ assert all(ti.region_id == region.id for ti in rows)
+
+
[email protected]("create_if_necessary", [False, "db", "memory"])
[email protected]("live_rows", [False, True])
+def test_cli_lookup_rejects_loop_scope_instead_of_creating_sentinel(
+ dag_maker, session, create_if_necessary, live_rows
+):
+ @task_group
+ def body():
+ BashOperator(task_id="terminal", bash_command="true")
+
+ with dag_maker(serialized=True):
+ group = create_loop(body, max_iterations=2)
+ dr = dag_maker.create_dagrun()
+ region = DynamicRegion(dag_id=dr.dag_id, run_id=dr.run_id,
node_id=group.group_id)
+ session.add(region)
+ session.flush()
+ for ti in dr.task_instances:
+ ti.region_id, ti.region_index = region.id, 0
+ if not live_rows:
+ session.execute(delete(TaskInstance).where(TaskInstance.dag_id ==
dr.dag_id))
+ session.commit()
+ serialized = dag_maker.serialized_dag
+
+ with pytest.raises(ValueError, match="loop.*scope"):
+ task_command._get_ti(
+ serialized.get_task(group.terminal_task_id),
+ -1,
+ logical_date_or_run_id=dr.run_id,
+ create_if_necessary=create_if_necessary,
+ session=session,
+ )
+
+ assert not session.scalars(
+ select(TaskInstance).where(TaskInstance.dag_id == dr.dag_id,
TaskInstance.region_id != region.id)
+ ).all()
+
+
[email protected]("create_if_necessary", ["db", "memory"])
[email protected]("regional_expansion", [False, True])
+def test_cli_missing_mapped_slot_is_created_ad_hoc_in_existing_expansion(
+ dag_maker, session, create_if_necessary, regional_expansion
+):
+ with dag_maker(serialized=True):
+ PythonOperator.partial(task_id="mapped",
python_callable=str).expand(op_args=[[1], [2]])
+ dr = dag_maker.create_dagrun()
+ if regional_expansion:
+ region =
session.scalars(select(DynamicRegion).where(DynamicRegion.dag_id ==
dr.dag_id)).one()
+ expected_region_id = region.id
+ session.execute(delete(TaskInstance).where(TaskInstance.dag_id ==
dr.dag_id))
+ else:
+ expected_region_id = SENTINEL_REGION_ID
+ for ti in dr.task_instances:
+ ti.region_id = SENTINEL_REGION_ID
+ session.flush()
+ session.execute(delete(DynamicRegion).where(DynamicRegion.dag_id ==
dr.dag_id))
+ session.execute(
+ delete(TaskInstance).where(TaskInstance.dag_id == dr.dag_id,
TaskInstance.region_index == 0)
+ )
+ session.commit()
+
+ ti, _ = task_command._get_ti(
+ dag_maker.serialized_dag.get_task("mapped"),
+ 0,
+ logical_date_or_run_id=dr.run_id,
+ create_if_necessary=create_if_necessary,
+ session=session,
+ )
+
+ assert (ti.region_id, ti.region_index) == (expected_region_id, 0)
+ assert session.scalars(select(DynamicRegion).where(DynamicRegion.dag_id ==
dr.dag_id)).all() == (
+ [region] if regional_expansion else []
+ )
+
+
[email protected]("create_if_necessary", ["db", "memory"])
+def test_cli_xcom_mapped_slot_is_created_ad_hoc_beside_placeholder(dag_maker,
session, create_if_necessary):
+ with dag_maker(serialized=True):
+ upstream = PythonOperator(task_id="upstream", python_callable=lambda:
[1, 2, 3])
+ PythonOperator.partial(task_id="mapped",
python_callable=str).expand(op_args=upstream.output)
+ dr = dag_maker.create_dagrun()
+ placeholder = session.scalars(
+ select(TaskInstance).where(TaskInstance.dag_id == dr.dag_id,
TaskInstance.task_id == "mapped")
+ ).all()
+ assert placeholder
+
+ ti, _ = task_command._get_ti(
+ dag_maker.serialized_dag.get_task("mapped"),
+ 2,
+ logical_date_or_run_id=dr.run_id,
+ create_if_necessary=create_if_necessary,
+ session=session,
+ )
+
+ assert ti.region_index == 2
+ assert ti.task_id == "mapped"
Review Comment:
The name promises the slot is created beside the placeholder, but these two
asserts only check the values passed in. A `_get_ti` that minted a fresh region
or fell back to the sentinel would still pass. Could you assert `ti.region_id
== placeholder[0].region_id != SENTINEL_REGION_ID`, and for `"db"` that the
task still has exactly one `DynamicRegion`?
--
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]