This is an automated email from the ASF dual-hosted git repository.
shahar1 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new d21496f7581 Fetch existing task instances in instead of rebuilding
them in tests (#74146)
d21496f7581 is described below
commit d21496f7581ba53eba076218f372e679f734af4d
Author: Ash Berlin-Taylor <[email protected]>
AuthorDate: Sat Oct 3 13:19:15 2026 +0100
Fetch existing task instances in instead of rebuilding them in tests
(#74146)
This is a "simple" correctness/tests-not-doing-what-code-does type fix.
A few tests constructed a transient TaskInstance for coordinates
(task_id, run_id, dag_id, etc) that already have a row, then call
refresh_from_db() and rely on it adopting the persisted row. That works
only because refresh_from_db() looks the row up by (dag_id, run_id,
task_id, map_index) and copies its id onto the new object. The transient
object is otherwise a second, unsaved instance with a fresh UUID, so the
first merge or flush after a refresh that finds nothing would insert a
duplicate row.
Fetching the persisted instance from its DagRun states what the tests mean
and removes their dependency on that lookup, which a task attempt keyed by
its UUID cannot honour.
In short: this works today, but isn't good practice, and it breaks a
future change I'm making, so I've extracted this simple text fix to it's
own PR.
---
airflow-core/tests/unit/jobs/test_scheduler_job.py | 10 +++------
airflow-core/tests/unit/models/test_dagrun.py | 5 +----
.../tests/unit/models/test_taskinstance.py | 25 +++++++++-------------
3 files changed, 14 insertions(+), 26 deletions(-)
diff --git a/airflow-core/tests/unit/jobs/test_scheduler_job.py
b/airflow-core/tests/unit/jobs/test_scheduler_job.py
index 08d9588b8bc..6b487bf6347 100644
--- a/airflow-core/tests/unit/jobs/test_scheduler_job.py
+++ b/airflow-core/tests/unit/jobs/test_scheduler_job.py
@@ -1561,17 +1561,15 @@ class TestSchedulerJob:
task_id_1 = "dummy_task"
with dag_maker(dag_id=dag_id):
- task1 = EmptyOperator(task_id=task_id_1)
+ EmptyOperator(task_id=task_id_1)
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(scheduler_job,
executors=[self.null_exec])
session = settings.Session()
dr1 = dag_maker.create_dagrun(run_type=DagRunType.BACKFILL_JOB)
- dag_version = DagVersion.get_latest_version(dr1.dag_id)
- ti1 = create_task_instance(task1, run_id=dr1.run_id,
dag_version_id=dag_version.id)
- ti1.refresh_from_db()
+ ti1 = dr1.get_task_instance(task_id_1, session=session)
ti1.state = State.SCHEDULED
session.merge(ti1)
session.flush()
@@ -8248,9 +8246,7 @@ class TestSchedulerJob:
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(job=scheduler_job,
executors=[MockExecutor(do_update=False)])
- dag_version = DagVersion.get_latest_version(dag_id=dag.dag_id)
- ti = create_task_instance(task=task1, run_id=dr1_running.run_id,
dag_version_id=dag_version.id)
- ti.refresh_from_db()
+ ti = dr1_running.get_task_instance(task1.task_id, session=session)
ti.state = State.SUCCESS
session.merge(ti)
session.flush()
diff --git a/airflow-core/tests/unit/models/test_dagrun.py
b/airflow-core/tests/unit/models/test_dagrun.py
index e897a59f35e..47343eeefae 100644
--- a/airflow-core/tests/unit/models/test_dagrun.py
+++ b/airflow-core/tests/unit/models/test_dagrun.py
@@ -991,12 +991,9 @@ class TestDagRun:
run_type=DagRunType.SCHEDULED,
)
- prev_ti = TI(task, run_id=dag_run_1.run_id,
dag_version_id=dag_run_1.created_dag_version_id)
- prev_ti.refresh_from_db(session=session)
+ prev_ti = dag_run_1.get_task_instance(task.task_id, session=session)
prev_ti.set_state(prev_ti_state, session=session)
session.flush()
- ti = TI(task, run_id=dag_run_2.run_id,
dag_version_id=dag_run_1.created_dag_version_id)
- ti.refresh_from_db(session=session)
decision =
dag_run_2.task_instance_scheduling_decisions(session=session)
schedulable_tis = [ti.task_id for ti in decision.schedulable_tis]
diff --git a/airflow-core/tests/unit/models/test_taskinstance.py
b/airflow-core/tests/unit/models/test_taskinstance.py
index c156341da6a..f88174f8b3b 100644
--- a/airflow-core/tests/unit/models/test_taskinstance.py
+++ b/airflow-core/tests/unit/models/test_taskinstance.py
@@ -1476,9 +1476,8 @@ class TestTaskInstance:
)
serialized_dag = SerializedDagModel.get(ti.task.dag.dag_id).dag
- ti_from_deserialized_task = TI(
- task=serialized_dag.get_task(ti.task_id), run_id=ti.run_id,
dag_version_id=ti.dag_version_id
- )
+ ti.task = serialized_dag.get_task(ti.task_id)
+ ti_from_deserialized_task = ti
assert ti_from_deserialized_task.try_number == 0
assert
ti_from_deserialized_task.check_and_change_state_before_execution()
@@ -1498,9 +1497,8 @@ class TestTaskInstance:
assert ti.external_executor_id == "apple"
serialized_dag = SerializedDagModel.get(ti.task.dag.dag_id).dag
- ti_from_deserialized_task = TI(
- task=serialized_dag.get_task(ti.task_id), run_id=ti.run_id,
dag_version_id=ti.dag_version_id
- )
+ ti.task = serialized_dag.get_task(ti.task_id)
+ ti_from_deserialized_task = ti
assert ti_from_deserialized_task.try_number == 0
assert
ti_from_deserialized_task.check_and_change_state_before_execution(
@@ -1517,9 +1515,8 @@ class TestTaskInstance:
assert ti.external_executor_id is None
serialized_dag = SerializedDagModel.get(ti.task.dag.dag_id).dag
- ti_from_deserialized_task = TI(
- task=serialized_dag.get_task(ti.task_id), run_id=ti.run_id,
dag_version_id=ti.dag_version_id
- )
+ ti.task = serialized_dag.get_task(ti.task_id)
+ ti_from_deserialized_task = ti
assert ti_from_deserialized_task.try_number == 0
assert
ti_from_deserialized_task.check_and_change_state_before_execution(
@@ -1565,9 +1562,8 @@ class TestTaskInstance:
ti.state = State.RUNNING
serialized_dag = SerializedDagModel.get(ti.task.dag.dag_id).dag
- ti_from_deserialized_task = TI(
- task=serialized_dag.get_task(ti.task_id), run_id=ti.run_id,
dag_version_id=ti.dag_version_id
- )
+ ti.task = serialized_dag.get_task(ti.task_id)
+ ti_from_deserialized_task = ti
assert not
ti_from_deserialized_task.check_and_change_state_before_execution()
assert ti_from_deserialized_task.state == State.RUNNING
@@ -1582,9 +1578,8 @@ class TestTaskInstance:
ti.state = State.FAILED
serialized_dag = SerializedDagModel.get(ti.task.dag.dag_id).dag
- ti_from_deserialized_task = TI(
- task=serialized_dag.get_task(ti.task_id), run_id=ti.run_id,
dag_version_id=ti.dag_version_id
- )
+ ti.task = serialized_dag.get_task(ti.task_id)
+ ti_from_deserialized_task = ti
assert not
ti_from_deserialized_task.check_and_change_state_before_execution()
assert ti_from_deserialized_task.state == State.FAILED