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

Reply via email to