Skip to content

Commit bc13269

Browse files
vatsrahul1001ashb
andauthored
Correctly pre-allocate external_exeuctor_id with multiple executors. (#67388) (#67458)
(cherry picked from commit 2def802) Co-authored-by: Ash Berlin-Taylor <ash@apache.org>
1 parent af90d7c commit bc13269

2 files changed

Lines changed: 47 additions & 12 deletions

File tree

airflow-core/src/airflow/jobs/scheduler_job_runner.py

Lines changed: 18 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,22 @@
3333
from itertools import groupby
3434
from typing import TYPE_CHECKING, Any, cast
3535

36-
from sqlalchemy import CTE, and_, case, delete, exists, func, inspect, or_, select, text, tuple_, update
36+
from sqlalchemy import (
37+
CTE,
38+
Text,
39+
and_,
40+
case,
41+
cast as sql_cast,
42+
delete,
43+
exists,
44+
func,
45+
inspect,
46+
or_,
47+
select,
48+
text,
49+
tuple_,
50+
update,
51+
)
3752
from sqlalchemy.exc import DBAPIError, OperationalError
3853
from sqlalchemy.orm import joinedload, lazyload, load_only, make_transient, selectinload
3954
from sqlalchemy.sql import expression
@@ -941,9 +956,9 @@ def _executable_task_instances_to_queued(self, max_tis: int, session: Session) -
941956
opt_in_names.add(exc.name.module_path)
942957
whens = []
943958
if opt_in_names:
944-
whens.append((TI.executor.in_(opt_in_names), random_db_uuid()))
959+
whens.append((TI.executor.in_(opt_in_names), sql_cast(random_db_uuid(), Text)))
945960
if default_opts_in:
946-
whens.append((TI.executor.is_(None), random_db_uuid()))
961+
whens.append((TI.executor.is_(None), sql_cast(random_db_uuid(), Text)))
947962
if whens:
948963
queued_values["external_executor_id"] = case(*whens, else_=TI.external_executor_id)
949964

airflow-core/tests/unit/jobs/test_scheduler_job.py

Lines changed: 29 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -2788,31 +2788,51 @@ def test_executable_task_instances_to_queued_sets_external_executor_id(self, dag
27882788
dag_id = "SchedulerJobTest.test_executable_sets_external_executor_id"
27892789
session = settings.Session()
27902790
with dag_maker(dag_id=dag_id, start_date=DEFAULT_DATE, session=session):
2791-
EmptyOperator(task_id="dummy")
2791+
EmptyOperator(task_id="a_task_pre_assign")
2792+
EmptyOperator(task_id="b_task_regular")
27922793

27932794
class PreAssigningExecutor(MockExecutor):
27942795
pre_assigns_external_executor_id = True
2796+
mock_module_path = "mock.pre_assigning.executor"
2797+
mock_alias = "pre_assigning_executor"
27952798

2796-
scheduler_job = Job()
2797-
self.job_runner = SchedulerJobRunner(job=scheduler_job, executors=[PreAssigningExecutor()])
2799+
regular_exec = MockExecutor()
2800+
assert regular_exec.pre_assigns_external_executor_id is False, "Pre-condition"
2801+
2802+
pre_assigning_exec = PreAssigningExecutor()
2803+
2804+
self.job_runner = SchedulerJobRunner(job=Job(), executors=(regular_exec, pre_assigning_exec))
27982805

27992806
dr = dag_maker.create_dagrun()
2800-
ti = dr.get_task_instance("dummy", session)
2801-
ti.state = State.SCHEDULED
2802-
session.merge(ti)
2807+
ti_pre_assign = dr.get_task_instance("a_task_pre_assign", session)
2808+
ti_regular = dr.get_task_instance("b_task_regular", session)
2809+
2810+
ti_regular.state = State.SCHEDULED
2811+
ti_regular.executor = regular_exec.name.module_path
2812+
ti_pre_assign.state = State.SCHEDULED
2813+
ti_pre_assign.executor = pre_assigning_exec.name.module_path
28032814
session.flush()
28042815

28052816
returned_tis = self.job_runner._executable_task_instances_to_queued(max_tis=32, session=session)
2817+
returned_tis.sort(key=lambda ti: ti.task_id)
2818+
2819+
assert len(returned_tis) == 2
28062820

2807-
assert len(returned_tis) == 1
28082821
# In-memory object (post make_transient) should carry the UUID
2822+
assert returned_tis[0].id == ti_pre_assign.id
28092823
assert returned_tis[0].external_executor_id is not None
2810-
UUID(returned_tis[0].external_executor_id)
2824+
assert UUID(returned_tis[0].external_executor_id), "is valid uuid"
28112825

28122826
# DB row should also have it (the whole point — survives a crash)
2813-
db_value = session.scalar(select(TaskInstance.external_executor_id).where(TaskInstance.id == ti.id))
2827+
db_value = session.scalar(
2828+
select(TaskInstance.external_executor_id).where(TaskInstance.id == ti_pre_assign.id)
2829+
)
28142830
assert db_value == returned_tis[0].external_executor_id
28152831

2832+
# In mixed-executor mode, only TIs routed to a pre-assigning executor get an external_executor_id.
2833+
assert returned_tis[1].id == ti_regular.id
2834+
assert returned_tis[1].external_executor_id is None
2835+
28162836
session.rollback()
28172837

28182838
@pytest.mark.parametrize("state", [State.FAILED, State.SUCCESS])

0 commit comments

Comments
 (0)