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)
