This is an automated email from the ASF dual-hosted git repository.

ashb pushed a commit to branch store-historic-ti-ownership-data
in repository https://gitbox.apache.org/repos/asf/airflow.git

commit 4652e57b0df645146f86f649566889fe245738c4
Author: Ash Berlin-Taylor <[email protected]>
AuthorDate: Mon Oct 5 21:03:44 2026 +0100

    fixup! Reduced diffs now we added with_loader_options
---
 .../core_api/services/public/task_instances.py     |   4 +-
 airflow-core/src/airflow/models/taskinstance.py    |   8 +-
 airflow-core/src/airflow/models/xcom.py            |   4 +-
 .../core_api/routes/public/test_extra_links.py     |   1 -
 .../core_api/routes/public/test_hitl.py            |   1 -
 .../core_api/routes/public/test_task_instances.py  |   6 +-
 .../versions/head/test_task_instances.py           |  44 +++------
 .../versions/v2026_10_30/test_task_instances.py    |   6 +-
 .../unit/cli/commands/test_partition_command.py    |   7 +-
 airflow-core/tests/unit/jobs/test_scheduler_job.py | 105 ++++-----------------
 airflow-core/tests/unit/models/test_cleartasks.py  |   3 +-
 airflow-core/tests/unit/models/test_dag.py         |  11 +--
 airflow-core/tests/unit/models/test_dagrun.py      |  12 +--
 .../tests/unit/models/test_mappedoperator.py       |   2 -
 .../tests/unit/models/test_taskinstance.py         |  30 ++----
 airflow-core/tests/unit/models/test_trigger.py     |   8 +-
 .../src/tests_common/test_utils/mapping.py         |   4 -
 17 files changed, 51 insertions(+), 205 deletions(-)

diff --git 
a/airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py
 
b/airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py
index c3f8d375849..4ad484c5a98 100644
--- 
a/airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py
+++ 
b/airflow-core/src/airflow/api_fastapi/core_api/services/public/task_instances.py
@@ -484,9 +484,7 @@ class 
BulkTaskInstanceService(BulkService[BulkTaskInstanceBody]):
         # Filter at database level using exact tuple matching instead of 
fetching all combinations
         # and filtering in Python
         task_keys_list = list(task_keys)
-        query = select(TI).where(
-            tuple_(TI.dag_id, TI.run_id, TI.task_id, 
TI.map_index).in_(task_keys_list),
-        )
+        query = select(TI).where(tuple_(TI.dag_id, TI.run_id, TI.task_id, 
TI.map_index).in_(task_keys_list))
 
         task_instances = self.session.scalars(query).all()
         task_instances_map = {
diff --git a/airflow-core/src/airflow/models/taskinstance.py 
b/airflow-core/src/airflow/models/taskinstance.py
index d801fd7e747..1aef66e04c3 100644
--- a/airflow-core/src/airflow/models/taskinstance.py
+++ b/airflow-core/src/airflow/models/taskinstance.py
@@ -2365,13 +2365,7 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload):
 
     @staticmethod
     def filter_for_tis(tis: Iterable[TaskInstance | TaskInstanceKey]) -> 
ColumnElement[bool] | None:
-        """Return SQLAlchemy filter to query the current attempt of the 
selected task instances."""
-        if (keys_filter := TaskInstance._build_keys_filter(tis)) is None:
-            return None
-        return and_(TaskInstance.working_set.is_(True), keys_filter)
-
-    @staticmethod
-    def _build_keys_filter(tis: Iterable[TaskInstance | TaskInstanceKey]) -> 
ColumnElement[bool] | None:
+        """Return SQLAlchemy filter to query selected task instances."""
         # DictKeys type, (what we often pass here from the scheduler) is not 
directly indexable :(
         # Or it might be a generator, but we need to be able to iterate over 
it more than once
         tis = list(tis)
diff --git a/airflow-core/src/airflow/models/xcom.py 
b/airflow-core/src/airflow/models/xcom.py
index 38184defa15..71e683e9295 100644
--- a/airflow-core/src/airflow/models/xcom.py
+++ b/airflow-core/src/airflow/models/xcom.py
@@ -585,9 +585,7 @@ def select_producers(
     from airflow.models.taskinstance import TaskInstance
 
     query = select(TaskInstance.id)
-    if try_number is None:
-        query = query.where(TaskInstance.working_set.is_(True))
-    else:
+    if try_number is not None:
         query = query.where(TaskInstance.try_number == try_number)
     for column, value in ((TaskInstance.dag_id, dag_ids), 
(TaskInstance.task_id, task_ids)):
         if is_container(value):
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_extra_links.py
 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_extra_links.py
index 8bf0034b983..1e55311fd19 100644
--- 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_extra_links.py
+++ 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_extra_links.py
@@ -446,7 +446,6 @@ class TestGetExtraLinks:
                 TaskInstance.run_id == self.dag_run_id,
                 TaskInstance.task_id == self.task_single_link,
                 TaskInstance.map_index == -1,
-                TaskInstance.working_set.is_(True),
             )
         )
         assert ti is not None
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_hitl.py 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_hitl.py
index 00152a18e29..2f8286c6609 100644
--- a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_hitl.py
+++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_hitl.py
@@ -364,7 +364,6 @@ class TestUpdateHITLDetailEndpoint:
                     TIModel.dag_id == sample_ti.dag_id,
                     TIModel.task_id == sample_ti.task_id,
                     TIModel.run_id == sample_ti.run_id,
-                    TIModel.working_set.is_(True),
                 )
             )
             assert current.id != sample_ti.id
diff --git 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
index 14a12e48fa8..b0daa900181 100644
--- 
a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
+++ 
b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_task_instances.py
@@ -183,7 +183,6 @@ class TestTaskInstanceEndpoint:
                         TaskInstance.task_id == ti.task_id,
                         TaskInstance.run_id == ti.run_id,
                         TaskInstance.map_index == ti.map_index,
-                        TaskInstance.working_set.is_(True),
                     )
                 )
                 assert current.id != ti.id
@@ -2957,7 +2956,7 @@ class TestGetTaskInstanceTry(TestTaskInstanceEndpoint):
             session.flush()
             session.add(RTIF(ti, render_templates=False))
         session.commit()
-        tis = 
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))).all()
+        tis = session.scalars(select(TaskInstance)).all()
         # Record the task instance history
         from airflow.models.taskinstance import clear_task_instances
 
@@ -4707,7 +4706,7 @@ class TestGetTaskInstanceTries(TestTaskInstanceEndpoint):
         self.create_task_instances(
             session=session, task_instances=[{"state": State.FAILED}], 
with_ti_history=True
         )
-        ti = 
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))).one()
+        ti = session.scalars(select(TaskInstance)).one()
         ti.state = State.UP_FOR_RETRY
         session.commit()
 
@@ -7688,7 +7687,6 @@ class TestPatchTaskGroup(TestTaskInstanceEndpoint):
                 TaskInstance.dag_id == self.DAG_ID,
                 TaskInstance.run_id == self.RUN_ID,
                 TaskInstance.task_id.in_(downstream_task_ids),
-                TaskInstance.working_set.is_(True),
             )
         ).all()
         assert {ti.task_id for ti in downstream_tis} == 
set(downstream_task_ids)
diff --git 
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
 
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
index 1f1ca7481a6..6c8039f84fd 100644
--- 
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
+++ 
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_task_instances.py
@@ -1437,11 +1437,7 @@ class TestTIUpdateState:
 
         assert response.status_code == (204 if matching_worker else 409)
         session.expunge_all()
-        current = session.scalar(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.task_id == "stopped_restart")
-        )
+        current = 
session.scalar(select(TaskInstance).where(TaskInstance.task_id == 
"stopped_restart"))
         if not matching_worker:
             assert (current.id, current.try_number, current.state) == (old_id, 
3, State.RESTARTING)
             return
@@ -1540,11 +1536,7 @@ class TestTIUpdateState:
                     SchedulerJobRunner.process_executor_events(executor, None, 
DBDagBag(), session)
                     session.commit()
             session.expunge_all()
-            current = session.scalar(
-                select(TaskInstance)
-                .where(TaskInstance.working_set.is_(True))
-                .where(TaskInstance.task_id == "restart_reports")
-            )
+            current = 
session.scalar(select(TaskInstance).where(TaskInstance.task_id == 
"restart_reports"))
             if definition == "broken_task":
                 assert (current.id, current.try_number, current.state, 
current.max_tries) == (
                     old_id,
@@ -2173,7 +2165,7 @@ class TestTIUpdateState:
 
             session.expire_all()
 
-            tis = 
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))).all()
+            tis = session.scalars(select(TaskInstance)).all()
             assert len(tis) == 1
 
             assert tis[0].state == TaskInstanceState.DEFERRED
@@ -2302,7 +2294,7 @@ class TestTIUpdateState:
 
         session.expire_all()
 
-        tis = 
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))).all()
+        tis = session.scalars(select(TaskInstance)).all()
         assert len(tis) == 1
         assert tis[0].state == TaskInstanceState.UP_FOR_RESCHEDULE
         assert tis[0].next_method is None
@@ -2456,9 +2448,7 @@ class TestTIUpdateState:
         assert response.text == ""
 
         ti = session.scalar(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .filter_by(task_id=ti.task_id, run_id=ti.run_id, dag_id=ti.dag_id)
+            select(TaskInstance).filter_by(task_id=ti.task_id, 
run_id=ti.run_id, dag_id=ti.dag_id)
         )
         # ti = session.get(TaskInstance, ti.id)
         assert ti.state == State.UP_FOR_RETRY
@@ -2514,9 +2504,7 @@ class TestTIUpdateState:
         assert client.patch(f"/execution/task-instances/{old_id}/state", 
json=payload).status_code == 204
         session.expunge_all()
         replacement = session.scalar(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.task_id == "retired_attempt_report")
+            select(TaskInstance).where(TaskInstance.task_id == 
"retired_attempt_report")
         )
         replacement.state = replacement_state
         replacement.external_executor_id = "replacement-worker"
@@ -2569,9 +2557,7 @@ class TestTIUpdateState:
         assert response.status_code == 204
 
         ti = session.scalar(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .filter_by(task_id=ti.task_id, run_id=ti.run_id, dag_id=ti.dag_id)
+            select(TaskInstance).filter_by(task_id=ti.task_id, 
run_id=ti.run_id, dag_id=ti.dag_id)
         )
         assert ti.state == State.UP_FOR_RETRY
         assert ti.retry_delay_override == 42.5
@@ -2649,9 +2635,7 @@ class TestTIUpdateState:
 
         session.expire_all()
         ti = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .filter_by(task_id=ti.task_id, run_id=ti.run_id, dag_id=ti.dag_id)
+            select(TaskInstance).filter_by(task_id=ti.task_id, 
run_id=ti.run_id, dag_id=ti.dag_id)
         ).one()
         assert ti.rendered_map_index is None
 
@@ -2686,9 +2670,7 @@ class TestTIUpdateState:
         assert response.status_code == 204
 
         ti = session.scalar(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .filter_by(task_id=ti.task_id, run_id=ti.run_id, dag_id=ti.dag_id)
+            select(TaskInstance).filter_by(task_id=ti.task_id, 
run_id=ti.run_id, dag_id=ti.dag_id)
         )
         assert ti.state == State.UP_FOR_RETRY
         assert ti.retry_delay_override is None
@@ -2733,9 +2715,7 @@ class TestTIUpdateState:
 
         session.expire_all()
         ti = session.scalar(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .filter_by(task_id=ti.task_id, run_id=ti.run_id, dag_id=ti.dag_id)
+            select(TaskInstance).filter_by(task_id=ti.task_id, 
run_id=ti.run_id, dag_id=ti.dag_id)
         )
         assert ti.state == State.RUNNING
         assert ti.retry_delay_override is None
@@ -2775,9 +2755,7 @@ class TestTIUpdateState:
         if target_state == State.UP_FOR_RETRY:
             # Retry creates a new TI ID, so we need to fetch by unique key
             ti = session.scalar(
-                select(TaskInstance)
-                .where(TaskInstance.working_set.is_(True))
-                .filter_by(task_id=ti.task_id, run_id=ti.run_id, 
dag_id=ti.dag_id)
+                select(TaskInstance).filter_by(task_id=ti.task_id, 
run_id=ti.run_id, dag_id=ti.dag_id)
             )
         else:
             session.expire_all()
diff --git 
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/v2026_10_30/test_task_instances.py
 
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/v2026_10_30/test_task_instances.py
index eecffc755a4..1dd49b684ee 100644
--- 
a/airflow-core/tests/unit/api_fastapi/execution_api/versions/v2026_10_30/test_task_instances.py
+++ 
b/airflow-core/tests/unit/api_fastapi/execution_api/versions/v2026_10_30/test_task_instances.py
@@ -55,11 +55,7 @@ def test_retired_state_report_response_by_version(
 
     assert response.status_code == expected_status
     session.expunge_all()
-    replacement = session.scalar(
-        select(TaskInstance).where(
-            TaskInstance.task_id == "retired_state_report", 
TaskInstance.working_set.is_(True)
-        )
-    )
+    replacement = 
session.scalar(select(TaskInstance).where(TaskInstance.task_id == 
"retired_state_report"))
     assert replacement.id != old_id
     assert (replacement.try_number, replacement.state) == (2, 
State.UP_FOR_RETRY)
 
diff --git a/airflow-core/tests/unit/cli/commands/test_partition_command.py 
b/airflow-core/tests/unit/cli/commands/test_partition_command.py
index 9e02e2fa7bc..7892c5bd74d 100644
--- a/airflow-core/tests/unit/cli/commands/test_partition_command.py
+++ b/airflow-core/tests/unit/cli/commands/test_partition_command.py
@@ -102,11 +102,7 @@ def _set_tis_state(run_id: str, state: TaskInstanceState) 
-> None:
 
 def _get_tis(run_id: str) -> list[TaskInstance]:
     with create_session() as session:
-        return list(
-            session.scalars(
-                select(TaskInstance).where(TaskInstance.run_id == run_id, 
TaskInstance.working_set.is_(True))
-            )
-        )
+        return 
list(session.scalars(select(TaskInstance).where(TaskInstance.run_id == run_id)))
 
 
 @pytest.mark.usefixtures("setup_partitioned_runs")
@@ -1338,7 +1334,6 @@ class TestPartitionsClear:
                     select(TaskInstance).where(
                         TaskInstance.dag_id == dag_id_target,
                         TaskInstance.run_id == shared_run_id,
-                        TaskInstance.working_set.is_(True),
                     )
                 )
             )
diff --git a/airflow-core/tests/unit/jobs/test_scheduler_job.py 
b/airflow-core/tests/unit/jobs/test_scheduler_job.py
index 102aa314599..e2539c64c0f 100644
--- a/airflow-core/tests/unit/jobs/test_scheduler_job.py
+++ b/airflow-core/tests/unit/jobs/test_scheduler_job.py
@@ -693,7 +693,6 @@ class TestSchedulerJob:
             select(TaskInstance).where(
                 TaskInstance.dag_id == dag_id,
                 TaskInstance.task_id == task_id,
-                TaskInstance.working_set.is_(True),
             )
         )
         assert replacement.state is None, "Replacement should be ready to 
schedule after termination"
@@ -2580,7 +2579,6 @@ class TestSchedulerJob:
             session.scalar(
                 select(func.count())
                 .select_from(TaskInstance)
-                .where(TaskInstance.working_set.is_(True))
                 .where(TaskInstance.dag_id == dag_id, TaskInstance.state == 
State.SCHEDULED)
             )
             == 1
@@ -2589,7 +2587,6 @@ class TestSchedulerJob:
             session.scalar(
                 select(func.count())
                 .select_from(TaskInstance)
-                .where(TaskInstance.working_set.is_(True))
                 .where(TaskInstance.dag_id == dag_id, TaskInstance.state == 
State.QUEUED)
             )
             == 1
@@ -2708,7 +2705,7 @@ class TestSchedulerJob:
         assert queued_runs["run_3"] == 2
 
         session.commit()
-        
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))).all()
+        session.scalars(select(TaskInstance)).all()
 
         # now we still have max tis running so no more will be queued
         queued_tis = self.job_runner._select_task_instances_to_queue(
@@ -5028,9 +5025,7 @@ class TestSchedulerJob:
 
         # Verify the task instance was created
         initial_tis = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.dag_id == dag_id, TaskInstance.task_id == 
"dummy")
+            select(TaskInstance).where(TaskInstance.dag_id == dag_id, 
TaskInstance.task_id == "dummy")
         ).all()
         assert len(initial_tis) == 1
 
@@ -5057,9 +5052,7 @@ class TestSchedulerJob:
 
         # Verify no new task instances were created for the removed task in 
the new dagrun
         new_tis = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(
+            select(TaskInstance).where(
                 TaskInstance.dag_id == dag_id,
                 TaskInstance.task_id == "dummy",
                 TaskInstance.run_id == "test_run_2",
@@ -5143,16 +5136,7 @@ class TestSchedulerJob:
             run_job(scheduler_job, execute_callable=self.job_runner._execute)
 
             # zero tasks ran
-            assert (
-                len(
-                    session.scalars(
-                        select(TaskInstance)
-                        .where(TaskInstance.working_set.is_(True))
-                        .where(TaskInstance.dag_id == dag_id)
-                    ).all()
-                )
-                == 0
-            )
+            assert 
len(session.scalars(select(TaskInstance).where(TaskInstance.dag_id == 
dag_id)).all()) == 0
             session.commit()
             assert self.null_exec.sorted_tasks == []
 
@@ -5171,16 +5155,7 @@ class TestSchedulerJob:
                 run_after=data_interval_end,
             )
             # one task "ran"
-            assert (
-                len(
-                    session.scalars(
-                        select(TaskInstance)
-                        .where(TaskInstance.working_set.is_(True))
-                        .where(TaskInstance.dag_id == dag_id)
-                    ).all()
-                )
-                == 1
-            )
+            assert 
len(session.scalars(select(TaskInstance).where(TaskInstance.dag_id == 
dag_id)).all()) == 1
             session.commit()
 
             scheduler_job = Job()
@@ -5189,16 +5164,7 @@ class TestSchedulerJob:
             run_job(scheduler_job, execute_callable=self.job_runner._execute)
 
             # still one task
-            assert (
-                len(
-                    session.scalars(
-                        select(TaskInstance)
-                        .where(TaskInstance.working_set.is_(True))
-                        .where(TaskInstance.dag_id == dag_id)
-                    ).all()
-                )
-                == 1
-            )
+            assert 
len(session.scalars(select(TaskInstance).where(TaskInstance.dag_id == 
dag_id)).all()) == 1
             session.commit()
             assert self.null_exec.sorted_tasks == []
 
@@ -5230,14 +5196,10 @@ class TestSchedulerJob:
 
         session = settings.Session()
         ti1s = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.dag_id == dag_id, TaskInstance.task_id == 
"dummy1")
+            select(TaskInstance).where(TaskInstance.dag_id == dag_id, 
TaskInstance.task_id == "dummy1")
         ).all()
         ti2s = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.dag_id == dag_id, TaskInstance.task_id == 
"dummy2")
+            select(TaskInstance).where(TaskInstance.dag_id == dag_id, 
TaskInstance.task_id == "dummy2")
         ).all()
 
         # With catchup=True, future task start dates are respected
@@ -5273,14 +5235,10 @@ class TestSchedulerJob:
 
         session = settings.Session()
         ti1s = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.dag_id == dag_id, TaskInstance.task_id == 
"dummy1")
+            select(TaskInstance).where(TaskInstance.dag_id == dag_id, 
TaskInstance.task_id == "dummy1")
         ).all()
         ti2s = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.dag_id == dag_id, TaskInstance.task_id == 
"dummy2")
+            select(TaskInstance).where(TaskInstance.dag_id == dag_id, 
TaskInstance.task_id == "dummy2")
         ).all()
 
         # With catchup=False, future task start dates are ignored
@@ -5318,16 +5276,7 @@ class TestSchedulerJob:
         # zero tasks ran
         dag_id = "test_start_date_scheduling"
         session = settings.Session()
-        assert (
-            len(
-                session.scalars(
-                    select(TaskInstance)
-                    .where(TaskInstance.working_set.is_(True))
-                    .where(TaskInstance.dag_id == dag_id)
-                ).all()
-            )
-            == 0
-        )
+        assert 
len(session.scalars(select(TaskInstance).where(TaskInstance.dag_id == 
dag_id)).all()) == 0
 
     def test_scheduler_verify_pool_full(self, dag_maker, mock_executor):
         """
@@ -5538,23 +5487,17 @@ class TestSchedulerJob:
         assert len(task_instances_list) == 2
 
         ti0 = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.task_id == 
"test_scheduler_verify_priority_and_slots_t0")
+            select(TaskInstance).where(TaskInstance.task_id == 
"test_scheduler_verify_priority_and_slots_t0")
         ).first()
         assert ti0.state == State.SCHEDULED
 
         ti1 = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.task_id == 
"test_scheduler_verify_priority_and_slots_t1")
+            select(TaskInstance).where(TaskInstance.task_id == 
"test_scheduler_verify_priority_and_slots_t1")
         ).first()
         assert ti1.state == State.QUEUED
 
         ti2 = session.scalars(
-            select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
-            .where(TaskInstance.task_id == 
"test_scheduler_verify_priority_and_slots_t2")
+            select(TaskInstance).where(TaskInstance.task_id == 
"test_scheduler_verify_priority_and_slots_t2")
         ).first()
         assert ti2.state == State.QUEUED
 
@@ -5752,9 +5695,7 @@ class TestSchedulerJob:
         do_schedule()
         with create_session() as session:
             ti = session.scalars(
-                select(TaskInstance)
-                .where(TaskInstance.working_set.is_(True))
-                .where(
+                select(TaskInstance).where(
                     TaskInstance.dag_id == "test_retry_still_in_executor",
                     TaskInstance.task_id == "test_retry_handling_op",
                 )
@@ -5777,7 +5718,6 @@ class TestSchedulerJob:
             session.expire_all()
             return session.scalar(
                 select(TaskInstance).where(
-                    TaskInstance.working_set.is_(True),
                     TaskInstance.dag_id == "test_retry_still_in_executor",
                     TaskInstance.task_id == "test_retry_handling_op",
                 )
@@ -5902,7 +5842,6 @@ class TestSchedulerJob:
             select(TaskInstance).where(
                 TaskInstance.dag_id == dag_id,
                 TaskInstance.task_id == task_id,
-                TaskInstance.working_set.is_(True),
             )
         )
 
@@ -7427,7 +7366,6 @@ class TestSchedulerJob:
         def complete_one_dagrun():
             ti = session.scalars(
                 select(TaskInstance)
-                .where(TaskInstance.working_set.is_(True))
                 .join(TaskInstance.dag_run)
                 .where(TaskInstance.state != State.SUCCESS)
                 .order_by(DagRun.logical_date)
@@ -8416,7 +8354,6 @@ class TestSchedulerJob:
                 session.scalar(
                     select(func.count())
                     .select_from(TaskInstance)
-                    .where(TaskInstance.working_set.is_(True))
                     .where(TaskInstance.state == State.SCHEDULED)
                 )
                 == 1
@@ -8470,7 +8407,6 @@ class TestSchedulerJob:
                 session.scalar(
                     select(func.count())
                     .select_from(TaskInstance)
-                    .where(TaskInstance.working_set.is_(True))
                     .where(TaskInstance.state == State.SCHEDULED)
                 )
                 == 1
@@ -8524,7 +8460,6 @@ class TestSchedulerJob:
                 session.scalar(
                     select(func.count())
                     .select_from(TaskInstance)
-                    .where(TaskInstance.working_set.is_(True))
                     .where(TaskInstance.state == State.SCHEDULED)
                 )
                 == 1
@@ -8585,7 +8520,6 @@ class TestSchedulerJob:
                 session.scalar(
                     select(func.count())
                     .select_from(TaskInstance)
-                    .where(TaskInstance.working_set.is_(True))
                     .where(TaskInstance.state == State.SCHEDULED)
                 )
                 == 2
@@ -8637,10 +8571,7 @@ class TestSchedulerJob:
         session.expunge_all()
         assert (
             session.scalar(
-                select(func.count())
-                .select_from(TaskInstance)
-                .where(TaskInstance.working_set.is_(True))
-                .where(TaskInstance.state == State.SCHEDULED)
+                
select(func.count()).select_from(TaskInstance).where(TaskInstance.state == 
State.SCHEDULED)
             )
             == 2
         )
@@ -9210,7 +9141,7 @@ class TestSchedulerJob:
         self.job_runner._schedule_dag_run(dr, session)
         session.expunge_all()
         with create_session() as session:
-            tis = 
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))).all()
+            tis = session.scalars(select(TaskInstance)).all()
 
         dags = [entry.dag for entry in 
self.job_runner.scheduler_dag_bag._dags.values()]
         assert [dag.dag_id for dag in dags] == ["test_only_empty_tasks"]
@@ -9238,7 +9169,7 @@ class TestSchedulerJob:
         self.job_runner._schedule_dag_run(dr, session)
         session.expunge_all()
         with create_session() as session:
-            tis = 
session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(True))).all()
+            tis = session.scalars(select(TaskInstance)).all()
 
         assert len(tis) == 6
         assert {
diff --git a/airflow-core/tests/unit/models/test_cleartasks.py 
b/airflow-core/tests/unit/models/test_cleartasks.py
index 9a872d3ed9c..c66e0c26579 100644
--- a/airflow-core/tests/unit/models/test_cleartasks.py
+++ b/airflow-core/tests/unit/models/test_cleartasks.py
@@ -866,7 +866,6 @@ class TestClearTasks:
                     TI.task_id == old_ti.task_id,
                     TI.map_index == old_ti.map_index,
                     TI.run_id == old_ti.run_id,
-                    TI.working_set.is_(True),
                 )
             )
 
@@ -1110,7 +1109,7 @@ class TestClearTasks:
         session.commit()
 
         dr_after = session.scalar(select(DagRun).where(DagRun.dag_id == 
dag_id))
-        ti_after = session.scalar(select(TI).where(TI.dag_id == dag_id, 
TI.working_set.is_(True)))
+        ti_after = session.scalar(select(TI).where(TI.dag_id == dag_id))
         assert dr_after.created_dag_version_id == new_dag_version.id
         assert ti_after.dag_version_id == dr_after.created_dag_version_id, (
             "the run and its task instance must end up on the same version"
diff --git a/airflow-core/tests/unit/models/test_dag.py 
b/airflow-core/tests/unit/models/test_dag.py
index 38dddaa352c..9af2717a66b 100644
--- a/airflow-core/tests/unit/models/test_dag.py
+++ b/airflow-core/tests/unit/models/test_dag.py
@@ -2166,9 +2166,7 @@ my_postgres_conn:
             session=session,
         )
 
-        task_instances = session.scalars(
-            select(TI).where(TI.dag_id == dag_id, TI.working_set.is_(True))
-        ).all()
+        task_instances = session.scalars(select(TI).where(TI.dag_id == 
dag_id)).all()
 
         assert len(task_instances) == 1
         task_instance: TI = task_instances[0]
@@ -3615,7 +3613,6 @@ def test_set_task_instance_state(run_id, session, 
dag_maker):
                 TI.dag_id == dag.dag_id,
                 TI.task_id == task.task_id,
                 TI.run_id == dagrun.run_id,
-                TI.working_set.is_(True),
             )
         )
 
@@ -3695,11 +3692,7 @@ def test_set_task_instance_state_mapped(dag_maker, 
session):
 
     ti_query = (
         select(TI.task_id, TI.map_index, TI.run_id, TI.state)
-        .where(
-            TI.dag_id == dag.dag_id,
-            TI.task_id.in_([task_id, "downstream"]),
-            TI.working_set.is_(True),
-        )
+        .where(TI.dag_id == dag.dag_id, TI.task_id.in_([task_id, 
"downstream"]))
         .order_by(TI.run_id, TI.task_id, TI.map_index)
     )
 
diff --git a/airflow-core/tests/unit/models/test_dagrun.py 
b/airflow-core/tests/unit/models/test_dagrun.py
index ecf204a16b2..926524c9cbe 100644
--- a/airflow-core/tests/unit/models/test_dagrun.py
+++ b/airflow-core/tests/unit/models/test_dagrun.py
@@ -3498,7 +3498,6 @@ def 
test_mapped_task_rerun_with_different_length_of_args(session, dag_maker, rer
         TI.run_id == dr.run_id,
         TI.task_id == "mapped_print_value",
         TI.state == TaskInstanceState.SUCCESS,
-        TI.working_set.is_(True),
     )
     success_tis = session.execute(query).all()
     assert len(success_tis) == rerun_length
@@ -3551,12 +3550,7 @@ def 
test_mapped_task_length_reduction_rerun_downstream_not_deadlocked(session, d
 
     mapped_states = session.execute(
         select(TI.map_index, TI.state)
-        .where(
-            TI.task_id == "work",
-            TI.dag_id == dr.dag_id,
-            TI.run_id == dr.run_id,
-            TI.working_set.is_(True),
-        )
+        .where(TI.task_id == "work", TI.dag_id == dr.dag_id, TI.run_id == 
dr.run_id)
         .order_by(TI.map_index)
     ).all()
     assert mapped_states == [
@@ -3842,9 +3836,7 @@ def 
test_clearing_task_and_moving_from_non_mapped_to_mapped(dag_maker, session):
     dr1: DagRun = dag_maker.create_dagrun(run_type=DagRunType.SCHEDULED)
     ti = dr1.get_task_instances()[0]
     ti = session.scalar(
-        select(TaskInstance)
-        .where(TaskInstance.working_set.is_(True))
-        .where(
+        select(TaskInstance).where(
             TaskInstance.dag_id == ti.dag_id,
             TaskInstance.task_id == ti.task_id,
             TaskInstance.run_id == ti.run_id,
diff --git a/airflow-core/tests/unit/models/test_mappedoperator.py 
b/airflow-core/tests/unit/models/test_mappedoperator.py
index 2fc29802b10..7834cf39fee 100644
--- a/airflow-core/tests/unit/models/test_mappedoperator.py
+++ b/airflow-core/tests/unit/models/test_mappedoperator.py
@@ -193,7 +193,6 @@ def test_expand_mapped_task_failed_state_in_db(dag_maker, 
session):
     indices = session.execute(
         select(TaskInstance.map_index, TaskInstance.state)
         .where(
-            TaskInstance.working_set.is_(True),
             TaskInstance.task_id == mapped.task_id,
             TaskInstance.dag_id == mapped.dag_id,
             TaskInstance.run_id == dr.run_id,
@@ -208,7 +207,6 @@ def test_expand_mapped_task_failed_state_in_db(dag_maker, 
session):
     indices = session.execute(
         select(TaskInstance.map_index, TaskInstance.state, 
TaskInstance.dag_version_id)
         .where(
-            TaskInstance.working_set.is_(True),
             TaskInstance.task_id == mapped.task_id,
             TaskInstance.dag_id == mapped.dag_id,
             TaskInstance.run_id == dr.run_id,
diff --git a/airflow-core/tests/unit/models/test_taskinstance.py 
b/airflow-core/tests/unit/models/test_taskinstance.py
index 4b56674d5f2..80083b4d80d 100644
--- a/airflow-core/tests/unit/models/test_taskinstance.py
+++ b/airflow-core/tests/unit/models/test_taskinstance.py
@@ -2768,16 +2768,11 @@ class TestTaskInstance:
         try_id = ti.id
         with pytest.raises(AirflowException):
             run_task_instance(ti, task)
-        ti = 
session.scalar(select(TaskInstance).where(TaskInstance.working_set.is_(True)))
+        ti = session.scalar(select(TaskInstance))
         # the ti.id should be different from the previous one
         assert ti.id != try_id
         assert ti.state == State.UP_FOR_RETRY
-        assert (
-            session.scalar(
-                
select(func.count()).select_from(TaskInstance).where(TaskInstance.working_set.is_(True))
-            )
-            == 1
-        )
+        assert session.scalar(select(func.count()).select_from(TaskInstance)) 
== 1
         tih = session.scalars(
             select(TaskInstance)
             .where(TaskInstance.working_set.is_(None))
@@ -3591,9 +3586,7 @@ class TestTaskInstanceRelationships:
         session.merge(ti)
         session.commit()
 
-        loaded_ti = session.scalar(
-            
select(TaskInstance).where(TaskInstance.working_set.is_(True)).where(TaskInstance.id
 == ti.id)
-        )
+        loaded_ti = session.scalar(select(TaskInstance).where(TaskInstance.id 
== ti.id))
 
         with pytest.raises(InvalidRequestError):
             getattr(loaded_ti, attr)
@@ -3945,7 +3938,6 @@ class TestMappedTaskInstanceReceiveValue:
 
         tis = session.scalars(
             select(TaskInstance)
-            .where(TaskInstance.working_set.is_(True))
             .where(
                 TaskInstance.dag_id == dag.dag_id,
                 TaskInstance.task_id == "show",
@@ -4147,12 +4139,7 @@ def test_taskinstance_with_note(create_task_instance, 
session):
     session.delete(ti)
     session.commit()
 
-    assert (
-        session.scalar(
-            
select(TaskInstance).where(TaskInstance.working_set.is_(True)).where(TaskInstance.id
 == ti.id)
-        )
-        is None
-    )
+    assert session.scalar(select(TaskInstance).where(TaskInstance.id == 
ti.id)) is None
     assert 
session.scalar(select(TaskInstanceNote).where(TaskInstanceNote.ti_id == ti.id)) 
is None
 
 
@@ -4160,7 +4147,7 @@ def 
test__refresh_from_db_should_not_increment_try_number(dag_maker, session):
     with dag_maker():
         BashOperator(task_id="hello", bash_command="hi")
     dag_maker.create_dagrun(state="success")
-    ti = 
session.scalar(select(TaskInstance).where(TaskInstance.working_set.is_(True)))
+    ti = session.scalar(select(TaskInstance))
     session.get(TaskInstance, ti.id).try_number += 1
     session.commit()
     assert ti.task_id == "hello"  # just to confirm...
@@ -4182,11 +4169,7 @@ def 
test_delete_dagversion_restricted_when_taskinstance_exists(dag_maker, sessio
     version = session.scalar(select(DagVersion).where(DagVersion.dag_id == 
dag.dag_id))
     assert version is not None
 
-    ti = session.scalars(
-        select(TaskInstance)
-        .where(TaskInstance.working_set.is_(True))
-        .where(TaskInstance.dag_version_id == version.id)
-    ).first()
+    ti = 
session.scalars(select(TaskInstance).where(TaskInstance.dag_version_id == 
version.id)).first()
     assert ti is not None
     if retired:
         ti.state = TaskInstanceState.SUCCESS
@@ -4899,7 +4882,6 @@ def 
test_task_instance_repr_does_not_raise_for_deferred_columns(dag_maker, sessi
     session.expunge_all()
     reloaded = session.scalar(
         select(TaskInstance)
-        .where(TaskInstance.working_set.is_(True))
         .where(TaskInstance.id == ti_id)
         .options(load_only(TaskInstance.dag_id, TaskInstance.task_id, 
TaskInstance.run_id))
     )
diff --git a/airflow-core/tests/unit/models/test_trigger.py 
b/airflow-core/tests/unit/models/test_trigger.py
index 97fa53a7133..dd24708e108 100644
--- a/airflow-core/tests/unit/models/test_trigger.py
+++ b/airflow-core/tests/unit/models/test_trigger.py
@@ -356,7 +356,7 @@ def test_submit_failure(session, create_task_instance):
     # Call submit_event
     Trigger.submit_failure(trigger.id, session=session)
     # Check that the task instance is now scheduled to fail
-    updated_task_instance = 
session.scalar(select(TaskInstance).where(TaskInstance.working_set.is_(True)))
+    updated_task_instance = session.scalar(select(TaskInstance))
     assert updated_task_instance.state == State.SCHEDULED
     assert updated_task_instance.next_method == "__fail__"
 
@@ -399,7 +399,7 @@ def test_submit_event_task_end(mock_utcnow, session, 
create_task_instance, event
 
     # now for the real test
     # first check initial state
-    ti: TaskInstance = 
session.scalar(select(TaskInstance).where(TaskInstance.working_set.is_(True)))
+    ti: TaskInstance = session.scalar(select(TaskInstance))
     assert ti.state == "deferred"
     assert get_xcoms(ti) == []
 
@@ -412,7 +412,7 @@ def test_submit_event_task_end(mock_utcnow, session, 
create_task_instance, event
     # commit changes made by submit event and expire all cache to read from db.
     session.flush()
     # Check that the task instance is now correct
-    ti = 
session.scalar(select(TaskInstance).where(TaskInstance.working_set.is_(True)))
+    ti = session.scalar(select(TaskInstance))
     assert ti.state == expected
     assert ti.next_kwargs is None
     assert ti.end_date == now
@@ -491,7 +491,7 @@ def test_submit_event_task_end_failed_respects_retries(
     Trigger.submit_event(trigger.id, TaskFailedEvent(), session=session)
     session.flush()
 
-    ti = 
session.scalar(select(TaskInstance).where(TaskInstance.working_set.is_(True)))
+    ti = session.scalar(select(TaskInstance))
     assert ti.state == expected_state
 
     mock_send.assert_called_once()
diff --git a/devel-common/src/tests_common/test_utils/mapping.py 
b/devel-common/src/tests_common/test_utils/mapping.py
index 1c787ffa099..bc75032884c 100644
--- a/devel-common/src/tests_common/test_utils/mapping.py
+++ b/devel-common/src/tests_common/test_utils/mapping.py
@@ -92,8 +92,6 @@ def expand_mapped_task_instances(
         .order_by(TaskInstance.map_index)
         .limit(1)
     )
-    if AIRFLOW_V_3_4_PLUS:
-        query = query.where(TaskInstance.working_set.is_(True))
     ti = session.scalars(query).one()
     ti.task = mapped
     return ti.expand_mapped_task(session=session)
@@ -112,8 +110,6 @@ def expand_mapped_task(
         TaskInstance.run_id == run_id,
         TaskInstance.map_index == -1,
     )
-    if AIRFLOW_V_3_4_PLUS:
-        query = query.where(TaskInstance.working_set.is_(True))
     upstream_ti = session.scalars(query).one()
     push_mapped_length(upstream_ti, list(range(length)), session=session)
     expand_mapped_task_instances(mapped, run_id, session=session)

Reply via email to