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 24d5836dc4c7fb0d1f6d8cd3509ca7b59d1830dd Author: Ash Berlin-Taylor <[email protected]> AuthorDate: Mon Oct 5 17:02:32 2026 +0100 fixup! Keep retired task attempts and their data under the attempt UUID --- airflow-core/src/airflow/api/common/delete_dag.py | 1 - airflow-core/src/airflow/api/common/mark_tasks.py | 3 - .../airflow/api_fastapi/common/parameters/misc.py | 1 - .../core_api/routes/public/extra_links.py | 20 ++-- .../api_fastapi/core_api/routes/public/hitl.py | 4 +- .../api_fastapi/core_api/routes/public/log.py | 1 + .../core_api/routes/public/task_instances.py | 18 ++- .../core_api/routes/public/task_state_store.py | 1 - .../airflow/api_fastapi/core_api/routes/ui/dags.py | 1 - .../api_fastapi/core_api/routes/ui/dashboard.py | 2 +- .../airflow/api_fastapi/core_api/routes/ui/grid.py | 2 - .../core_api/services/public/dag_run.py | 5 +- .../core_api/services/public/task_instances.py | 6 +- .../execution_api/routes/task_instances.py | 36 +++--- .../api_fastapi/execution_api/routes/xcoms.py | 9 +- .../airflow/api_fastapi/execution_api/security.py | 6 +- .../src/airflow/cli/commands/dag_command.py | 2 - .../src/airflow/jobs/scheduler_job_runner.py | 29 ++--- airflow-core/src/airflow/models/dagrun.py | 20 +--- airflow-core/src/airflow/models/pool.py | 11 +- .../src/airflow/models/renderedtifields.py | 8 +- airflow-core/src/airflow/models/taskinstance.py | 98 +++++++++++++--- airflow-core/src/airflow/models/trigger.py | 4 - airflow-core/src/airflow/models/xcom.py | 6 +- .../src/airflow/serialization/definitions/dag.py | 7 +- .../ti_deps/deps/mapped_task_upstream_dep.py | 1 - .../src/airflow/ti_deps/deps/trigger_rule_dep.py | 3 - .../core_api/routes/public/test_task_instances.py | 20 +++- .../api_fastapi/execution_api/test_security.py | 6 +- .../versions/head/test_task_instances.py | 30 +++-- airflow-core/tests/unit/jobs/test_scheduler_job.py | 22 +++- airflow-core/tests/unit/models/test_cleartasks.py | 37 +++++- airflow-core/tests/unit/models/test_dagrun.py | 10 +- airflow-core/tests/unit/models/test_task_data.py | 5 +- .../tests/unit/models/test_taskinstance.py | 124 +++++++++++++++++++-- airflow-core/tests/unit/models/test_trigger.py | 1 + airflow-core/tests/unit/utils/test_db_cleanup.py | 12 +- 37 files changed, 388 insertions(+), 184 deletions(-) diff --git a/airflow-core/src/airflow/api/common/delete_dag.py b/airflow-core/src/airflow/api/common/delete_dag.py index 6473d7dfaba..5df54993c08 100644 --- a/airflow-core/src/airflow/api/common/delete_dag.py +++ b/airflow-core/src/airflow/api/common/delete_dag.py @@ -57,7 +57,6 @@ def delete_dag(dag_id: str, keep_records_in_log: bool = True, *, session: Sessio select(models.TaskInstance.state) .where( models.TaskInstance.dag_id == dag_id, - models.TaskInstance.working_set.is_(True), models.TaskInstance.state == TaskInstanceState.RUNNING, ) .limit(1) diff --git a/airflow-core/src/airflow/api/common/mark_tasks.py b/airflow-core/src/airflow/api/common/mark_tasks.py index 6ba61686aec..31c81db0864 100644 --- a/airflow-core/src/airflow/api/common/mark_tasks.py +++ b/airflow-core/src/airflow/api/common/mark_tasks.py @@ -113,7 +113,6 @@ def get_all_dag_task_query( ): """Get all tasks of the main dag that will be affected by a state change.""" qry_dag = select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag.dag_id, TaskInstance.run_id.in_(run_ids), ) @@ -262,7 +261,6 @@ def _set_dag_run_terminal_state( running_tis: list[TaskInstance] = list( session.scalars( select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag.dag_id, TaskInstance.run_id == run_id, TaskInstance.task_id.in_(task_ids), @@ -285,7 +283,6 @@ def _set_dag_run_terminal_state( pending_tis: list[TaskInstance] = list( session.scalars( select(TaskInstance).filter( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag.dag_id, TaskInstance.run_id == run_id, TaskInstance.task_id.in_(task_ids), diff --git a/airflow-core/src/airflow/api_fastapi/common/parameters/misc.py b/airflow-core/src/airflow/api_fastapi/common/parameters/misc.py index be10b0d1fa5..6a5b2a44bdc 100644 --- a/airflow-core/src/airflow/api_fastapi/common/parameters/misc.py +++ b/airflow-core/src/airflow/api_fastapi/common/parameters/misc.py @@ -63,7 +63,6 @@ class _PendingActionsFilter(BaseParam[bool]): .join(TaskInstance, HITLDetail.ti_id == TaskInstance.id) .where( HITLDetail.responded_at.is_(None), - TaskInstance.working_set.is_(True), TaskInstance.state.in_((TaskInstanceState.DEFERRED, TaskInstanceState.AWAITING_INPUT)), ) .where(TaskInstance.dag_id == DagModel.dag_id) diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/extra_links.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/extra_links.py index 3b0ef365f05..e80c4c28b8f 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/extra_links.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/extra_links.py @@ -73,17 +73,17 @@ def get_extra_links( dag_run = session.scalar(select(DagRun).where(DagRun.dag_id == dag_id, DagRun.run_id == dag_run_id)) - ti = session.scalar( - select(TaskInstance).where( - TaskInstance.dag_id == dag_id, - TaskInstance.run_id == dag_run_id, - TaskInstance.task_id == task_id, - TaskInstance.map_index == map_index, - TaskInstance.working_set.is_(True) - if try_number is None - else TaskInstance.try_number == try_number, - ) + query = select(TaskInstance).where( + TaskInstance.dag_id == dag_id, + TaskInstance.run_id == dag_run_id, + TaskInstance.task_id == task_id, + TaskInstance.map_index == map_index, ) + if try_number is not None: + query = query.where(TaskInstance.try_number == try_number).execution_options( + include_all_attempts=True + ) + ti = session.scalar(query) if not ti: raise HTTPException( diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/hitl.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/hitl.py index abf4c497506..4ef1c6fb132 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/hitl.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/hitl.py @@ -99,7 +99,8 @@ def _get_task_instance_with_hitl_detail( ) .options(joinedload(TI.hitl_detail), joinedload(TI.rendered_task_instance_fields)) ) - query = query.where(TI.working_set.is_(True) if try_number is None else TI.try_number == try_number) + if try_number is not None: + query = query.where(TI.try_number == try_number).execution_options(include_all_attempts=True) ti = session.scalar(query) if ti is None: @@ -351,7 +352,6 @@ def get_hitl_details( query = ( select(HITLDetailModel) .join(TI, HITLDetailModel.ti_id == TI.id) - .where(TI.working_set.is_(True)) .join(TI.dag_run) .options( joinedload(HITLDetailModel.task_instance).options( diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/log.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/log.py index 58927314410..40f6ad16923 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/log.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/log.py @@ -134,6 +134,7 @@ def get_log( .join(TaskInstance.dag_run) .options(joinedload(TaskInstance.trigger).joinedload(Trigger.triggerer_job)) .options(joinedload(TaskInstance.dag_model)) + .execution_options(include_all_attempts=True) ) ti = session.scalar(query) if ti is None: diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_instances.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_instances.py index a6c0940ec42..1b77e22f512 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_instances.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_instances.py @@ -146,7 +146,7 @@ def get_task_instance( """Get task instance.""" query = ( select(TI) - .where(TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id) + .where(TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id) .options(joinedload(TI.rendered_task_instance_fields)) .options(joinedload(TI.dag_version)) .options(joinedload(TI.dag_run).options(joinedload(DagRun.dag_model))) @@ -243,7 +243,6 @@ def get_mapped_task_instances( """Get list of mapped task instances.""" query = eager_load_task_instance_for_validation( select(TI).where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id, @@ -325,9 +324,7 @@ def get_task_instance_dependencies( map_index: int = -1, ) -> TaskDependencyCollectionResponse: """Get dependencies blocking task from getting scheduled.""" - query = select(TI).where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id - ) + query = select(TI).where(TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id) query = query.where(TI.map_index == map_index) result = session.execute(query).one_or_none() @@ -390,6 +387,7 @@ def get_task_instance_tries( ) .options(joinedload(TI.hitl_detail)) .order_by(TI.try_number) + .execution_options(include_all_attempts=True) ) task_instances = list(session.scalars(query)) @@ -442,7 +440,6 @@ def get_mapped_task_instance( query = ( select(TI) .where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id, @@ -571,7 +568,7 @@ def get_task_instances( """ use_cursor = cursor is not None dag_run = None - query = eager_load_task_instance_for_validation(select(TI).where(TI.working_set.is_(True))) + query = eager_load_task_instance_for_validation(select(TI)) if dag_run_id != "~": if dag_id == "~": raise HTTPException( @@ -771,7 +768,7 @@ def get_task_instances_batch( TI, ).set_value([body.order_by] if body.order_by else None) - query = eager_load_task_instance_for_validation(select(TI).where(TI.working_set.is_(True))) + query = eager_load_task_instance_for_validation(select(TI)) task_instance_select, total_entries = paginated_select( statement=query, filters=[ @@ -818,13 +815,15 @@ def get_task_instance_try_details( ) -> TaskInstanceHistoryResponse: """Get task instance details by try number.""" query = eager_load_task_instance_for_validation( - select(TI).where( + select(TI) + .where( TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id, TI.try_number == task_try_number, TI.map_index == map_index, ) + .execution_options(include_all_attempts=True) ) ti = session.scalar(query) if ti is None: @@ -1299,7 +1298,6 @@ def delete_task_instance( ) -> None: """Delete a task instance.""" query = select(TI).where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id, diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_state_store.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_state_store.py index df6a8b9daab..8c0e5be5ef3 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_state_store.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/task_state_store.py @@ -59,7 +59,6 @@ def _require_task_instance( ) -> None: """Raise 404 unless the task instance exists. ``map_index=None`` matches any map index.""" statement = select(TI.task_id).where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id, diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dags.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dags.py index b23bbcfc426..aa25889284c 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dags.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dags.py @@ -261,7 +261,6 @@ def get_dags( defaultload(HITLDetail.task_instance).joinedload(TaskInstance.rendered_task_instance_fields) ) .where( - TaskInstance.working_set.is_(True), HITLDetail.responded_at.is_(None), TaskInstance.state.in_((TaskInstanceState.DEFERRED, TaskInstanceState.AWAITING_INPUT)), ) diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dashboard.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dashboard.py index f962916313c..97b2589c76e 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dashboard.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dashboard.py @@ -108,7 +108,7 @@ def historical_metrics( ) task_instance_states, task_instances_are_lower_bounds = _compute_state_counts( TaskInstance, - [*dag_run_filters, TaskInstance.working_set.is_(True)], + dag_run_filters, session=session, join=TaskInstance.dag_run, null_label="no_status", diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py index fbfd20b70e6..2dc91fa546b 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/grid.py @@ -226,7 +226,6 @@ def get_dag_structure( select(TaskInstance.dag_version_id) .join(TaskInstance.dag_run) .where( - TaskInstance.working_set.is_(True), DagRun.id.in_(run_ids), ) .distinct() @@ -540,7 +539,6 @@ def get_grid_ti_summaries_stream( ) .outerjoin(DagVersion, TaskInstance.dag_version_id == DagVersion.id) .where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag_id, TaskInstance.run_id == run_id, ) diff --git a/airflow-core/src/airflow/api_fastapi/core_api/services/public/dag_run.py b/airflow-core/src/airflow/api_fastapi/core_api/services/public/dag_run.py index 84803c27e21..d91fe06a077 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/services/public/dag_run.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/services/public/dag_run.py @@ -112,7 +112,6 @@ def dry_run_clear_dag_run( existing_task_ids = set( session.scalars( select(TaskInstance.task_id).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag_id, TaskInstance.run_id == dag_run_id, ) @@ -121,9 +120,7 @@ def dry_run_clear_dag_run( new_task_ids = sorted(set(latest_dag.task_ids) - existing_task_ids) return [NewTaskResponse(task_id=task_id, task_display_name=task_id) for task_id in new_task_ids] - ti_query = eager_load_task_instance_for_validation( - select(TaskInstance).where(TaskInstance.working_set.is_(True)) - ) + ti_query = eager_load_task_instance_for_validation(select(TaskInstance)) ti_query = ti_query.where( TaskInstance.dag_id == dag_id, TaskInstance.run_id == dag_run_id, 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 b8a443e4a66..c3f8d375849 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 @@ -204,7 +204,7 @@ def _patch_ti_validate_request( query = ( select(TI) - .where(TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id) + .where(TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id == task_id) .options(joinedload(TI.rendered_task_instance_fields)) ) if map_index is not None: @@ -247,7 +247,6 @@ def _get_task_group_task_instances( query = ( select(TI) .where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == dag_run_id, TI.task_id.in_(task_ids), @@ -486,7 +485,6 @@ class BulkTaskInstanceService(BulkService[BulkTaskInstanceBody]): # and filtering in Python task_keys_list = list(task_keys) query = select(TI).where( - TI.working_set.is_(True), tuple_(TI.dag_id, TI.run_id, TI.task_id, TI.map_index).in_(task_keys_list), ) @@ -608,7 +606,6 @@ class BulkTaskInstanceService(BulkService[BulkTaskInstanceBody]): batch_task_instances = self.session.scalars( select(TI).where( - TI.working_set.is_(True), TI.dag_id.in_(all_dag_ids), TI.run_id.in_(all_run_ids), TI.task_id.in_(all_task_ids), @@ -697,7 +694,6 @@ class BulkTaskInstanceService(BulkService[BulkTaskInstanceBody]): batch_task_instances = self.session.scalars( select(TI).where( - TI.working_set.is_(True), TI.dag_id.in_(all_dag_ids), TI.run_id.in_(all_run_ids), TI.task_id.in_(all_task_ids), diff --git a/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py b/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py index b9a36103e3b..0ea58145387 100644 --- a/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py +++ b/airflow-core/src/airflow/api_fastapi/execution_api/routes/task_instances.py @@ -187,7 +187,7 @@ def ti_run( .select_from(TI) .join(DR, and_(TI.dag_id == DR.dag_id, TI.run_id == DR.run_id)) .join(DagModel, TI.dag_id == DagModel.dag_id) - .where(TI.id == task_instance_id, TI.working_set.is_(True)) + .where(TI.id == task_instance_id) .with_for_update(of=TI) ) try: @@ -404,7 +404,7 @@ def ti_update_state( select(TI) .where(TI.id == task_instance_id) .with_for_update(of=TI) - .execution_options(populate_existing=True) + .execution_options(populate_existing=True, include_all_attempts=True) ) if ti is None: raise HTTPException(status_code=404, detail={"reason": "not_found"}) @@ -446,6 +446,7 @@ def ti_update_state( .join(DagModel, TI.dag_id == DagModel.dag_id) .where(TI.id == task_instance_id) .with_for_update(of=TI) + .execution_options(include_all_attempts=True) ) try: ( @@ -518,7 +519,7 @@ def ti_update_state( ) if "rendered_map_index" in data: data["_rendered_map_index"] = data.pop("rendered_map_index") - query = update(TI).where(TI.working_set.is_(True), TI.id == task_instance_id).values(data) + query = update(TI).where(TI.id == task_instance_id).values(data) asset_callbacks: Sequence[Callable[[], None]] = () try: @@ -540,9 +541,7 @@ def ti_update_state( payload=ti_patch_payload, ) session.rollback() - ti = session.scalar( - select(TI).where(TI.id == task_instance_id, TI.working_set.is_(True)).with_for_update(of=TI) - ) + ti = session.scalar(select(TI).where(TI.id == task_instance_id).with_for_update(of=TI)) if session.bind is not None: query = TI.duration_expression_update(timezone.utcnow(), query, session.bind) query = query.values(state=(updated_state := TaskInstanceState.FAILED)) @@ -901,9 +900,7 @@ def ti_skip_downstream( now = timezone.utcnow() tasks = ti_patch_payload.tasks - query_result = session.execute( - select(TI.dag_id, TI.run_id).where(TI.working_set.is_(True), TI.id == task_instance_id) - ) + query_result = session.execute(select(TI.dag_id, TI.run_id).where(TI.id == task_instance_id)) row_result = query_result.fetchone() if row_result is None: raise HTTPException( @@ -935,7 +932,6 @@ def ti_skip_downstream( query = ( update(TI) .where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == run_id, or_(TI.task_id.in_(task_ids), tuple_(TI.task_id, TI.map_index).in_(ti_keys)), @@ -1058,7 +1054,6 @@ async def ti_heartbeat( await session.execute( update(TI) .where( - TI.working_set.is_(True), TI.id == task_instance_id, TI.state == TaskInstanceState.RUNNING, TI.hostname == ti_payload.hostname, @@ -1078,6 +1073,7 @@ async def ti_heartbeat( select(TI.state, TI.hostname, TI.pid, TI.working_set) .where(TI.id == task_instance_id) .with_for_update() + .execution_options(include_all_attempts=True) ) try: @@ -1123,9 +1119,7 @@ async def ti_heartbeat( # Update the last heartbeat time! await session.execute( - update(TI) - .where(TI.working_set.is_(True), TI.id == task_instance_id) - .values(last_heartbeat_at=timezone.utcnow()) + update(TI).where(TI.id == task_instance_id).values(last_heartbeat_at=timezone.utcnow()) ) log.debug("Heartbeat updated", state=previous_state) @@ -1161,7 +1155,9 @@ def ti_put_rtif( bind_contextvars(ti_id=str(task_instance_id)) log.info("Updating RenderedTaskInstanceFields", field_count=len(put_rtif_payload)) - task_instance = session.scalar(select(TI).where(TI.id == task_instance_id)) + task_instance = session.scalar( + select(TI).where(TI.id == task_instance_id).execution_options(include_all_attempts=True) + ) if task_instance is None or task_instance.working_set is None: # On retry/clear, the server regenerates the TI id. Return 410 for the stale id. _raise_ti_not_in_live_table(task_instance_id, archived_in_history=task_instance is not None) @@ -1274,7 +1270,7 @@ def get_task_instance_count( states: Annotated[list[str] | None, Query()] = None, ) -> int: """Get the count of task instances matching the given criteria.""" - query = select(func.count()).select_from(TI).where(TI.dag_id == dag_id, TI.working_set.is_(True)) + query = select(func.count()).select_from(TI).where(TI.dag_id == dag_id) if task_ids: query = query.where(TI.task_id.in_(task_ids)) @@ -1340,9 +1336,7 @@ async def get_previous_task_instance( select(TI) .join(DR, (TI.dag_id == DR.dag_id) & (TI.run_id == DR.run_id)) .options(contains_eager(TI.dag_run).load_only(DR.logical_date)) - .where( - TI.dag_id == dag_id, TI.task_id == task_id, TI.map_index == map_index, TI.working_set.is_(True) - ) + .where(TI.dag_id == dag_id, TI.task_id == task_id, TI.map_index == map_index) .order_by(DR.logical_date.desc()) ) @@ -1386,7 +1380,7 @@ def get_task_instance_states( """Get the states for Task Instances with the given criteria.""" run_id_task_state_map: dict[str, dict[str, Any]] = defaultdict(dict) - query = select(TI).where(TI.working_set.is_(True), TI.dag_id == dag_id) + query = select(TI).where(TI.dag_id == dag_id) if task_ids: query = query.where(TI.task_id.in_(task_ids)) @@ -1429,7 +1423,6 @@ async def get_task_instance_breadcrumbs( await session.execute( select(TI.task_id, TI.map_index, TI.state, TI.operator, TI.duration) .where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == run_id, TI.state.in_(TerminalTIState), @@ -1481,7 +1474,6 @@ def _get_group_tasks( # First get all task instances to get the task_id, map_index pairs group_tasks = session.scalars( select(TI).where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.task_id.in_(task.task_id for task in task_group.iter_tasks()), *([TI.logical_date.in_(logical_dates)] if logical_dates else []), diff --git a/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py b/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py index d8ab77c22b4..542f95e4cb4 100644 --- a/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py +++ b/airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py @@ -342,6 +342,7 @@ def get_xcom( TaskInstance.map_index == params.map_index, ) .limit(1) + .execution_options(include_all_attempts=True) ).first() if ( owner is not None @@ -553,10 +554,14 @@ def _find_writer_id( TaskInstance.task_id == task_id, TaskInstance.map_index == map_index, ) - own = session.scalar(select(TaskInstance.id).where(TaskInstance.id == attempt_id, *coordinates)) + own = session.scalar( + select(TaskInstance.id) + .where(TaskInstance.id == attempt_id, *coordinates) + .execution_options(include_all_attempts=True) + ) if own is not None: return own - return session.scalar(select(TaskInstance.id).where(TaskInstance.working_set.is_(True), *coordinates)) + return session.scalar(select(TaskInstance.id).where(*coordinates)) def _get_writer_id( diff --git a/airflow-core/src/airflow/api_fastapi/execution_api/security.py b/airflow-core/src/airflow/api_fastapi/execution_api/security.py index 1f10e1c04a5..806c7fff4f5 100644 --- a/airflow-core/src/airflow/api_fastapi/execution_api/security.py +++ b/airflow-core/src/airflow/api_fastapi/execution_api/security.py @@ -250,7 +250,11 @@ async def _require_live_attempt(token: TIToken, *, allow_callback: bool) -> None """ async with create_session_async() as session: attempt = ( - await session.execute(select(TaskInstance.working_set).where(TaskInstance.id == token.id)) + await session.execute( + select(TaskInstance.working_set) + .where(TaskInstance.id == token.id) + .execution_options(include_all_attempts=True) + ) ).one_or_none() if attempt is not None and attempt.working_set: return diff --git a/airflow-core/src/airflow/cli/commands/dag_command.py b/airflow-core/src/airflow/cli/commands/dag_command.py index d7ced15f906..1f96179bfbb 100644 --- a/airflow-core/src/airflow/cli/commands/dag_command.py +++ b/airflow-core/src/airflow/cli/commands/dag_command.py @@ -226,7 +226,6 @@ def _bulk_clear_runs( cleared = 0 for chunk_run_ids in chunks(run_ids, _RUN_CHUNK_SIZE): ti_query = select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag_id, TaskInstance.run_id.in_(chunk_run_ids), ) @@ -862,7 +861,6 @@ def dag_test(args, dag: DAG | None = None, *, session: Session = NEW_SESSION) -> if show_dagrun or imgcat or filename: tis = session.scalars( select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag.dag_id, TaskInstance.run_id == dr.run_id, ) diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py b/airflow-core/src/airflow/jobs/scheduler_job_runner.py index 80f2bd0592d..f69c2654bfb 100644 --- a/airflow-core/src/airflow/jobs/scheduler_job_runner.py +++ b/airflow-core/src/airflow/jobs/scheduler_job_runner.py @@ -277,7 +277,7 @@ class ConcurrencyMap: self.task_dagrun_concurrency_map.clear() query = session.execute( select(TI.dag_id, TI.task_id, TI.run_id, TI.state, func.count("*")) - .where(TI.working_set.is_(True), TI.state.in_(ACTIVE_STATES)) + .where(TI.state.in_(ACTIVE_STATES)) .group_by(TI.dag_id, TI.task_id, TI.run_id, TI.state) ) for dag_id, task_id, run_id, state, count in query: @@ -306,7 +306,7 @@ def _get_current_dr_task_concurrency(states: Iterable[TaskInstanceState]) -> Sub """Get the dag_run IDs and how many tasks are in the provided states for each one.""" return ( select(TI.dag_id, TI.run_id, func.count("*").label("task_per_dr_count")) - .where(TI.working_set.is_(True), TI.state.in_(states)) + .where(TI.state.in_(states)) .group_by(TI.dag_id, TI.run_id) .subquery() ) @@ -610,7 +610,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): session.execute( update(TI) - .where(TI.working_set.is_(True), TI.dag_id == dag_id, TI.state == TaskInstanceState.SCHEDULED) + .where(TI.dag_id == dag_id, TI.state == TaskInstanceState.SCHEDULED) .values(state=TaskInstanceState.FAILED) .execution_options(synchronize_session="fetch") ) @@ -1020,7 +1020,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): select(TI) .with_hint(TI, "USE INDEX (ti_state)", dialect_name="mysql") .join(TI.dag_run) - .where(DR.state == DagRunState.RUNNING, TI.working_set.is_(True)) + .where(DR.state == DagRunState.RUNNING) .join(TI.dag_model) .where(~DM.is_paused) .where(TI.state == TaskInstanceState.SCHEDULED) @@ -1084,7 +1084,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): & (TI.run_id == ranked_query.c.run_id) & (TI.map_index == ranked_query.c.map_index), ) - .where(ranked_query.c.row_num <= ranked_query.c.dr_max_active_tasks, TI.working_set.is_(True)) + .where(ranked_query.c.row_num <= ranked_query.c.dr_max_active_tasks) # Add the order_by columns from the ranked query for sqlite. .order_by( -ranked_query.c.priority_weight_for_ordering, @@ -1530,7 +1530,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): asset_loader, alias_loader = _eager_load_dag_run_for_validation() query = ( select(TI) - .where(TI.working_set.is_(True), TI.id.in_([key.id for key in tis_with_right_state])) + .where(TI.id.in_([key.id for key in tis_with_right_state])) .options(selectinload(TI.dag_model)) .options(asset_loader) .options(alias_loader) @@ -3142,7 +3142,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): dag_run.set_state(DagRunState.FAILED) unfinished_task_instances = session.scalars( select(TI) - .where(TI.working_set.is_(True), TI.dag_id == dag_run.dag_id) + .where(TI.dag_id == dag_run.dag_id) .where(TI.run_id == dag_run.run_id) .where(TI.state.in_(State.unfinished) | (TI.state.is_(None))) ).all() @@ -3274,7 +3274,6 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): session.execute( update(TI) .where( - TI.working_set.is_(True), TI.dag_id == dag_run.dag_id, TI.run_id == dag_run.run_id, TI.state.in_(State.unfinished), @@ -3329,7 +3328,6 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): """Query db for TIs that are stuck in queued.""" return session.scalars( select(TI).where( - TI.working_set.is_(True), TI.state == TaskInstanceState.QUEUED, TI.queued_dttm < (timezone.utcnow() - timedelta(seconds=self._task_queued_timeout)), TI.queued_by_job_id == self.job.id, @@ -3488,7 +3486,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): TaskInstance.queue, func.count(TaskInstance.task_id).label("count"), ) - .filter(TaskInstance.state.in_(metric_states), TaskInstance.working_set.is_(True)) + .filter(TaskInstance.state.in_(metric_states)) .group_by(TaskInstance.state, TaskInstance.dag_id, TaskInstance.task_id, TaskInstance.queue) ) all_states_metric = session.execute(stmt).all() @@ -3627,7 +3625,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): query = ( select(TI) .options(lazyload(TI.dag_run)) # avoids double join to dag_run - .where(TI.state.in_(State.adoptable_states), TI.working_set.is_(True)) + .where(TI.state.in_(State.adoptable_states)) .join(TI.queued_by_job) .where(Job.state.is_distinct_from(JobState.RUNNING)) .join(TI.dag_run) @@ -3706,7 +3704,6 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): result = session.execute( update(TI) .where( - TI.working_set.is_(True), TI.state == TaskInstanceState.DEFERRED, TI.trigger_timeout < timezone.utcnow(), ) @@ -3740,7 +3737,6 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): query = ( select(TI) .where( - TI.working_set.is_(True), TI.state == TaskInstanceState.AWAITING_INPUT, TI.trigger_timeout < now, ) @@ -3863,7 +3859,6 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): .with_hint(TI, "USE INDEX (ti_state)", dialect_name="mysql") .join(DM, TI.dag_id == DM.dag_id) .where( - TI.working_set.is_(True), TI.state.in_((TaskInstanceState.RUNNING, TaskInstanceState.RESTARTING)), TI.last_heartbeat_at < limit_dttm, ) @@ -4058,11 +4053,7 @@ class SchedulerJobRunner(BaseJobRunner, LoggingMixin): select(AssetWatcherModel.trigger_id).where(AssetWatcherModel.trigger_id == Trigger.id) ), ~exists(select(Callback.trigger_id).where(Callback.trigger_id == Trigger.id)), - ~exists( - select(TaskInstance.trigger_id).where( - TaskInstance.working_set.is_(True), TaskInstance.trigger_id == Trigger.id - ) - ), + ~exists(select(TaskInstance.trigger_id).where(TaskInstance.trigger_id == Trigger.id)), ) .execution_options(synchronize_session="fetch") ) diff --git a/airflow-core/src/airflow/models/dagrun.py b/airflow-core/src/airflow/models/dagrun.py index 12d967a0f21..09c62db8026 100644 --- a/airflow-core/src/airflow/models/dagrun.py +++ b/airflow-core/src/airflow/models/dagrun.py @@ -591,6 +591,7 @@ class DagRun(Base, LoggingMixin): TI.run_id == self.run_id, ) .limit(1) + .execution_options(include_all_attempts=True) ) return session.scalar(select_stmt) @@ -973,7 +974,6 @@ class DagRun(Base, LoggingMixin): select(TI) .options(joinedload(TI.dag_run)) .where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id == run_id, ) @@ -1102,9 +1102,7 @@ class DagRun(Base, LoggingMixin): :param session: Sqlalchemy ORM Session """ return session.scalars( - select(TI) - .where(TI.working_set.is_(True)) - .filter_by(dag_id=dag_id, run_id=dag_run_id, task_id=task_id, map_index=map_index) + select(TI).filter_by(dag_id=dag_id, run_id=dag_run_id, task_id=task_id, map_index=map_index) ).one_or_none() def get_dag(self) -> SerializedDAG: @@ -1788,7 +1786,7 @@ class DagRun(Base, LoggingMixin): # Check if any ti changed state tis_filter = TI.filter_for_tis(old_states) if tis_filter is not None: - fresh_tis = session.scalars(select(TI).where(TI.working_set.is_(True), tis_filter)).all() + fresh_tis = session.scalars(select(TI).where(tis_filter)).all() changed_tis = any(ti.state != old_states[ti.key] for ti in fresh_tis) return ready_tis, changed_tis, expansion_happened @@ -2181,7 +2179,6 @@ class DagRun(Base, LoggingMixin): query = session.scalars( select(TI.map_index).where( - TI.working_set.is_(True), TI.dag_id == self.dag_id, TI.task_id == task.task_id, TI.run_id == self.run_id, @@ -2194,7 +2191,6 @@ class DagRun(Base, LoggingMixin): session.execute( update(TI) .where( - TI.working_set.is_(True), TI.dag_id == self.dag_id, TI.task_id == task.task_id, TI.run_id == self.run_id, @@ -2309,7 +2305,7 @@ class DagRun(Base, LoggingMixin): for id_chunk in schedulable_ti_ids_chunks: result = session.execute( update(TI) - .where(TI.working_set.is_(True), TI.id.in_(id_chunk), schedulable_state_clause) + .where(TI.id.in_(id_chunk), schedulable_state_clause) .values( state=TaskInstanceState.SCHEDULED, scheduled_dttm=timezone.utcnow(), @@ -2320,9 +2316,7 @@ class DagRun(Base, LoggingMixin): count += getattr(result, "rowcount", 0) if debug_try_number_check: rows = session.execute( - select(TI.id, TI.try_number, TI.state).where( - TI.working_set.is_(True), TI.id.in_(id_chunk) - ) + select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk)) ).all() rows_by_ti_id = { ti_id: (db_try_number, db_state) for ti_id, db_try_number, db_state in rows @@ -2359,7 +2353,7 @@ class DagRun(Base, LoggingMixin): for id_chunk in dummy_ti_ids_chunks: result = session.execute( update(TI) - .where(TI.working_set.is_(True), TI.id.in_(id_chunk), schedulable_state_clause) + .where(TI.id.in_(id_chunk), schedulable_state_clause) .values( state=TaskInstanceState.SUCCESS, start_date=timezone.utcnow(), @@ -2528,7 +2522,6 @@ def clear_partition_runs( chunk_tis = list( session.scalars( select(TI).where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id.in_(ti_buffer_run_ids), ) @@ -2577,7 +2570,6 @@ def clear_partition_runs( select(func.count()) .select_from(TI) .where( - TI.working_set.is_(True), TI.dag_id == dag_id, TI.run_id.in_(chunk), ) diff --git a/airflow-core/src/airflow/models/pool.py b/airflow-core/src/airflow/models/pool.py index c5b4d8dc428..3a2bc8f0922 100644 --- a/airflow-core/src/airflow/models/pool.py +++ b/airflow-core/src/airflow/models/pool.py @@ -222,7 +222,7 @@ class Pool(Base): } state_count_by_pool = session.execute( select(TaskInstance.pool, TaskInstance.state, func.sum(TaskInstance.pool_slots)) - .filter(TaskInstance.state.in_(allowed_execution_states), TaskInstance.working_set.is_(True)) + .filter(TaskInstance.state.in_(allowed_execution_states)) .group_by(TaskInstance.pool, TaskInstance.state) ) @@ -284,7 +284,7 @@ class Pool(Base): return int( session.scalar( select(func.sum(TaskInstance.pool_slots)) - .filter(TaskInstance.pool == self.pool, TaskInstance.working_set.is_(True)) + .filter(TaskInstance.pool == self.pool) .filter(TaskInstance.state.in_(occupied_states)) ) or 0 @@ -310,7 +310,7 @@ class Pool(Base): return int( session.scalar( select(func.sum(TaskInstance.pool_slots)) - .filter(TaskInstance.pool == self.pool, TaskInstance.working_set.is_(True)) + .filter(TaskInstance.pool == self.pool) .filter(TaskInstance.state == TaskInstanceState.RUNNING) ) or 0 @@ -329,7 +329,7 @@ class Pool(Base): return int( session.scalar( select(func.sum(TaskInstance.pool_slots)) - .filter(TaskInstance.pool == self.pool, TaskInstance.working_set.is_(True)) + .filter(TaskInstance.pool == self.pool) .filter(TaskInstance.state == TaskInstanceState.QUEUED) ) or 0 @@ -348,7 +348,7 @@ class Pool(Base): return int( session.scalar( select(func.sum(TaskInstance.pool_slots)) - .filter(TaskInstance.pool == self.pool, TaskInstance.working_set.is_(True)) + .filter(TaskInstance.pool == self.pool) .filter(TaskInstance.state == TaskInstanceState.SCHEDULED) ) or 0 @@ -367,7 +367,6 @@ class Pool(Base): return int( session.scalar( select(func.sum(TaskInstance.pool_slots)).where( - TaskInstance.working_set.is_(True), TaskInstance.pool == self.pool, TaskInstance.state == TaskInstanceState.DEFERRED, ) diff --git a/airflow-core/src/airflow/models/renderedtifields.py b/airflow-core/src/airflow/models/renderedtifields.py index 220d1bc78ae..cea27c1b446 100644 --- a/airflow-core/src/airflow/models/renderedtifields.py +++ b/airflow-core/src/airflow/models/renderedtifields.py @@ -208,11 +208,11 @@ class RenderedTaskInstanceFields(Base): ) .exists() ) - session.execute(delete(legacy).where(owns_legacy)) + session.execute(delete(legacy).where(owns_legacy).execution_options(include_all_attempts=True)) session.execute( delete(cls) .where(cls.task_instance_id.in_(producer_ids)) - .execution_options(synchronize_session="fetch") + .execution_options(synchronize_session="fetch", include_all_attempts=True) ) @staticmethod @@ -222,13 +222,15 @@ class RenderedTaskInstanceFields(Base): from airflow.models.taskinstance import TaskInstance return session.scalar( - select(TaskInstance.id).where( + select(TaskInstance.id) + .where( TaskInstance.dag_id == ti.dag_id, TaskInstance.task_id == ti.task_id, TaskInstance.run_id == ti.run_id, TaskInstance.map_index == ti.map_index, TaskInstance.try_number == ti.try_number, ) + .execution_options(include_all_attempts=True) ) def __init__(self, ti: TaskInstance, render_templates=True, rendered_fields=None): diff --git a/airflow-core/src/airflow/models/taskinstance.py b/airflow-core/src/airflow/models/taskinstance.py index 45d25f92545..ce187eb613d 100644 --- a/airflow-core/src/airflow/models/taskinstance.py +++ b/airflow-core/src/airflow/models/taskinstance.py @@ -52,6 +52,7 @@ from sqlalchemy import ( case, cast, delete, + event as sqlalchemy_event, extract, false, func, @@ -68,10 +69,19 @@ from sqlalchemy.dialects import postgresql from sqlalchemy.ext.associationproxy import association_proxy from sqlalchemy.ext.hybrid import hybrid_property from sqlalchemy.ext.mutable import MutableDict -from sqlalchemy.orm import Mapped, lazyload, mapped_column, reconstructor, relationship +from sqlalchemy.orm import ( + Mapped, + Session, + lazyload, + mapped_column, + reconstructor, + relationship, + with_loader_criteria, +) from sqlalchemy.orm.attributes import NO_VALUE, set_committed_value from sqlalchemy.orm.exc import DetachedInstanceError, ObjectDeletedError -from sqlalchemy.sql.elements import ColumnElement +from sqlalchemy.sql import visitors +from sqlalchemy.sql.elements import BindParameter, ColumnElement from airflow import settings from airflow._shared.observability.metrics import stats @@ -123,7 +133,7 @@ if TYPE_CHECKING: from typing import Literal from sqlalchemy.engine import Connection as SAConnection, Engine - from sqlalchemy.orm.session import Session + from sqlalchemy.orm import ORMExecuteState from sqlalchemy.sql import Update from sqlalchemy.sql.elements import ColumnElement @@ -367,7 +377,6 @@ def _pin_versionless_tis_to_run_version(dag_run: DagRun, dag_version_id: UUID, s session.execute( update(TaskInstance) .where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == dag_run.dag_id, TaskInstance.run_id == dag_run.run_id, TaskInstance.dag_version_id.is_(None), @@ -656,6 +665,24 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): A value of -1 in map_index represents any of: a TI without mapped tasks; a TI with mapped tasks that has yet to be expanded (state=pending); a TI with mapped tasks that expanded to an empty list (state=skipped). + + Every try of a task is its own row with its own UUID. Only the latest try is live (``working_set`` is + true); earlier tries are retired (``working_set`` is NULL) and kept as history. + + ORM queries see only live rows by default: a session hook adds ``working_set IS TRUE`` to every ORM + select, update and delete, including joins to this model. To include retired rows, set the execution + option on the statement:: + + session.scalars(select(TaskInstance).where(...).execution_options(include_all_attempts=True)) + + The default does not apply to: + + * primary key lookups (``Session.get``, ``merge``, ``refresh``), which return the row with that UUID + whether or not it is retired; + * relationship loads, which follow the join the relationship defines (``DagRun.task_instances`` is live, + ``DagRun.historical_task_instances`` is retired); + * Core statements on ``TaskInstance.__table__``, and an ``exists()`` that does not name this model in a + FROM clause. These see every row unless they filter ``working_set`` themselves. """ __tablename__ = "task_instance" @@ -982,7 +1009,6 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): ) -> TaskInstance | None: query = ( select(TaskInstance) - .where(TaskInstance.working_set.is_(True)) .options(lazyload(TaskInstance.dag_run)) # lazy load dag run to avoid locking it .filter_by( dag_id=dag_id, @@ -1019,10 +1045,14 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): :param keep_local_changes: Force all attributes to the values from the database if False (the default), or if True don't overwrite locally set attributes """ - query = select( - # Select the columns, not the ORM object, to bypass any session/ORM caching layer - *TaskInstance.__table__.columns - ).where(TaskInstance.id == self.id) + query = ( + select( + # Select the columns, not the ORM object, to bypass any session/ORM caching layer + *TaskInstance.__table__.columns + ) + .where(TaskInstance.id == self.id) + .execution_options(include_all_attempts=True) + ) if lock_for_update: query = query.with_for_update() @@ -1100,7 +1130,7 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): :param session: SQLAlchemy ORM Session :return: Was the state changed """ - if self.state == state: + if self.state == state or (self.working_set is None and inspect(self).has_identity): return False current_time = timezone.utcnow() @@ -1150,7 +1180,11 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): session: Session, ) -> None: """Delete every attempt, current and retired, of a task; all map indexes if ``map_index`` is None.""" - statement = delete(cls).where(cls.dag_id == dag_id, cls.run_id == run_id, cls.task_id == task_id) + statement = ( + delete(cls) + .where(cls.dag_id == dag_id, cls.run_id == run_id, cls.task_id == task_id) + .execution_options(include_all_attempts=True) + ) if map_index is not None: statement = statement.where(cls.map_index == map_index) session.execute(statement) @@ -1172,6 +1206,7 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): select(cls.map_index, func.max(cls.try_number)) .where(cls.dag_id == dag_id, cls.task_id == task_id, cls.run_id == run_id) .group_by(cls.map_index) + .execution_options(include_all_attempts=True) ) if map_indexes is not None: statement = statement.where(cls.map_index.in_(map_indexes)) @@ -1242,7 +1277,6 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): return True ti = select(func.count(TaskInstance.task_id)).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == self.dag_id, TaskInstance.task_id.in_(task.downstream_task_ids), TaskInstance.run_id == self.run_id, @@ -1950,7 +1984,7 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): with create_session() as session: session.execute( update(TaskInstance) - .where(TaskInstance.working_set.is_(True), TaskInstance.id == self.id) + .where(TaskInstance.id == self.id) .values(last_heartbeat_at=timezone.utcnow()) ) @@ -2320,7 +2354,6 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): select(func.count()) .select_from(TaskInstance) .where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == self.dag_id, TaskInstance.task_id == self.task_id, ) @@ -2525,7 +2558,6 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): state: str | None = None unmapped_ti: TaskInstance | None = session.scalars( select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == task.dag_id, TaskInstance.task_id == task.task_id, TaskInstance.run_id == run_id, @@ -2597,7 +2629,6 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): current_max_mapping = ( session.scalar( select(func.max(TaskInstance.map_index)).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == task.dag_id, TaskInstance.task_id == task.task_id, TaskInstance.run_id == run_id, @@ -2651,7 +2682,6 @@ class TaskInstance(Base, LoggingMixin, BaseWorkload): # Any (old) task instances with inapplicable indexes (>= the total # number we need) are set to "REMOVED". query = select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == task.dag_id, TaskInstance.task_id == task.task_id, TaskInstance.run_id == run_id, @@ -2956,6 +2986,40 @@ class TaskInstanceNote(Base): return prefix + f" TI ID: {self.ti_id}>" +_CURRENT_ATTEMPTS = with_loader_criteria( + TaskInstance, TaskInstance.working_set.is_(True), include_aliases=True, propagate_to_loaders=False +) + + +def _is_primary_key_lookup(statement) -> bool: + criteria: Sequence[Any] = getattr(statement, "_where_criteria", ()) + if len(criteria) != 1: + return False + binds = [node for node in visitors.iterate(criteria[0]) if isinstance(node, BindParameter)] + return bool(binds) and all(bind.key.startswith("pk_") for bind in binds) + + +@sqlalchemy_event.listens_for(Session, "do_orm_execute") +def _restrict_to_current_attempts(state: ORMExecuteState) -> None: + """ + Hide retired attempts from ORM queries over task instances unless ``include_all_attempts`` is set. + + Primary key lookups (``Session.get``, ``merge`` and ``refresh``) are exempt: asking for an attempt by + its UUID returns it whether or not it has been retired. + """ + if ( + state.is_column_load + or state.is_relationship_load + or state.execution_options.get("include_all_attempts") + ): + return + if not (state.is_select or state.is_update or state.is_delete): + return + if state.is_select and _is_primary_key_lookup(state.statement): + return + state.statement = state.statement.options(_CURRENT_ATTEMPTS) + + STATICA_HACK = True globals()["kcah_acitats"[::-1].upper()] = False if STATICA_HACK: # pragma: no cover diff --git a/airflow-core/src/airflow/models/trigger.py b/airflow-core/src/airflow/models/trigger.py index f86c82b8491..cbc9a67a109 100644 --- a/airflow-core/src/airflow/models/trigger.py +++ b/airflow-core/src/airflow/models/trigger.py @@ -245,7 +245,6 @@ class Trigger(Base): session.execute( update(TaskInstance) .where( - TaskInstance.working_set.is_(True), TaskInstance.state != TaskInstanceState.DEFERRED, TaskInstance.trigger_id.is_not(None), ) @@ -282,7 +281,6 @@ class Trigger(Base): # Resume deferred tasks for task_instance in session.scalars( select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.trigger_id == trigger_id, TaskInstance.state == TaskInstanceState.DEFERRED, ) @@ -323,7 +321,6 @@ class Trigger(Base): """ for task_instance in session.scalars( select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.trigger_id == trigger_id, TaskInstance.state == TaskInstanceState.DEFERRED, ) @@ -475,7 +472,6 @@ class Trigger(Base): .prefix_with("STRAIGHT_JOIN", dialect="mysql") .join(TaskInstance, cls.id == TaskInstance.trigger_id, isouter=False) .where( - TaskInstance.working_set.is_(True), or_(cls.triggerer_id.is_(None), cls.triggerer_id.not_in(alive_triggerer_ids)), ) .order_by(coalesce(TaskInstance.priority_weight, 0).desc(), cls.created_date), diff --git a/airflow-core/src/airflow/models/xcom.py b/airflow-core/src/airflow/models/xcom.py index 89c654200bb..38184defa15 100644 --- a/airflow-core/src/airflow/models/xcom.py +++ b/airflow-core/src/airflow/models/xcom.py @@ -214,8 +214,8 @@ class _XComOperations: if key is not None: v1_delete = v1_delete.where(legacy.c.key == key) v2_delete = v2_delete.where(XComModelV2.key == key) - session.execute(v1_delete) - session.execute(v2_delete.execution_options(synchronize_session="fetch")) + session.execute(v1_delete.execution_options(include_all_attempts=True)) + session.execute(v2_delete.execution_options(synchronize_session="fetch", include_all_attempts=True)) @classmethod @provide_session @@ -362,6 +362,8 @@ class _XComOperations: statement = statement.order_by(entity.logical_date.desc(), entity.timestamp.desc()) if limit: statement = statement.limit(limit) + if try_number is not None: + statement = statement.execution_options(include_all_attempts=True) return statement @staticmethod diff --git a/airflow-core/src/airflow/serialization/definitions/dag.py b/airflow-core/src/airflow/serialization/definitions/dag.py index 7fad6d46e29..0acb5a2faa3 100644 --- a/airflow-core/src/airflow/serialization/definitions/dag.py +++ b/airflow-core/src/airflow/serialization/definitions/dag.py @@ -540,7 +540,6 @@ class SerializedDAG: total_tasks = session.scalar( select(func.count(TaskInstance.task_id)).where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == self.dag_id, TaskInstance.state == TaskInstanceState.RUNNING, ) @@ -1068,10 +1067,10 @@ class SerializedDAG: TaskInstance.task_id, TaskInstance.run_id, TaskInstance.map_index, - ).where(TaskInstance.working_set.is_(True)) + ) tis_pk = tis_pk.join(TaskInstance.dag_run) else: - tis_full = select(TaskInstance).where(TaskInstance.working_set.is_(True)) + tis_full = select(TaskInstance) tis_full = tis_full.join(TaskInstance.dag_run) # Apply common filters @@ -1148,7 +1147,7 @@ class SerializedDAG: # We've been asked for objects, lets combine it all back in to a result set ti_filters = TaskInstance.filter_for_tis(result) if ti_filters is not None: - tis_final = select(TaskInstance).where(TaskInstance.working_set.is_(True), ti_filters) + tis_final = select(TaskInstance).where(ti_filters) return session.scalars(tis_final) elif exclude_task_ids is None: pass # Disable filter if not set. diff --git a/airflow-core/src/airflow/ti_deps/deps/mapped_task_upstream_dep.py b/airflow-core/src/airflow/ti_deps/deps/mapped_task_upstream_dep.py index afc0df74828..e1957b9a342 100644 --- a/airflow-core/src/airflow/ti_deps/deps/mapped_task_upstream_dep.py +++ b/airflow-core/src/airflow/ti_deps/deps/mapped_task_upstream_dep.py @@ -72,7 +72,6 @@ class MappedTaskUpstreamDep(BaseTIDep): mapped_dependency_tis = ( session.scalars( select(TaskInstance).where( - TaskInstance.working_set.is_(True), TaskInstance.task_id.in_(operator.task_id for operator in mapped_dependencies), TaskInstance.dag_id == ti.dag_id, TaskInstance.run_id == ti.run_id, diff --git a/airflow-core/src/airflow/ti_deps/deps/trigger_rule_dep.py b/airflow-core/src/airflow/ti_deps/deps/trigger_rule_dep.py index dd9b77ed44f..24b77000555 100644 --- a/airflow-core/src/airflow/ti_deps/deps/trigger_rule_dep.py +++ b/airflow-core/src/airflow/ti_deps/deps/trigger_rule_dep.py @@ -303,7 +303,6 @@ class TriggerRuleDep(BaseTIDep): task_id_counts = session.execute( select(TaskInstance.task_id, func.count(TaskInstance.task_id)) .where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == ti.dag_id, TaskInstance.run_id == ti.run_id, ) @@ -415,7 +414,6 @@ class TriggerRuleDep(BaseTIDep): for task_id, count in session.execute( select(TaskInstance.task_id, func.count(TaskInstance.task_id)) .where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == ti.dag_id, TaskInstance.run_id == ti.run_id, ) @@ -715,7 +713,6 @@ class TriggerRuleDep(BaseTIDep): session.scalar( select(func.count(TaskInstance.task_id)) .where( - TaskInstance.working_set.is_(True), TaskInstance.dag_id == ti.dag_id, TaskInstance.run_id == ti.run_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 f1c4317b435..14a12e48fa8 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 @@ -4580,7 +4580,11 @@ class TestGetTaskInstanceTries(TestTaskInstanceEndpoint): self.create_task_instances( session=session, task_instances=[{"state": State.SUCCESS}], with_ti_history=True ) - historical = session.scalar(select(TaskInstance).where(TaskInstance.working_set.is_(None))) + historical = session.scalar( + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .execution_options(include_all_attempts=True) + ) historical.dag_version_id = None session.commit() @@ -6439,11 +6443,13 @@ class TestBulkTaskInstances(TestTaskInstanceEndpoint): session.delete(next(ti for ti in current_tis if ti.map_index == 2)) session.flush() task_rows = session.scalars( - select(TaskInstance).where( + select(TaskInstance) + .where( TaskInstance.dag_id == self.DAG_ID, TaskInstance.run_id == self.RUN_ID, TaskInstance.task_id == self.TASK_ID, ) + .execution_options(include_all_attempts=True) ).all() ids_by_index = { map_index: {ti.id for ti in task_rows if ti.map_index == map_index} for map_index in (0, 1, 2) @@ -6457,7 +6463,11 @@ class TestBulkTaskInstances(TestTaskInstanceEndpoint): == 404 ) assert ( - session.scalar(select(TaskInstance.id).where(TaskInstance.id.in_(ids_by_index[2]))) + session.scalar( + select(TaskInstance.id) + .where(TaskInstance.id.in_(ids_by_index[2])) + .execution_options(include_all_attempts=True) + ) in ids_by_index[2] ) response = test_client.delete(f"{self.ENDPOINT_URL}/{self.TASK_ID}", params={"map_index": 0}) @@ -6481,11 +6491,13 @@ class TestBulkTaskInstances(TestTaskInstanceEndpoint): assert ( set( session.scalars( - select(TaskInstance.id).where( + select(TaskInstance.id) + .where( TaskInstance.dag_id == self.DAG_ID, TaskInstance.run_id == self.RUN_ID, TaskInstance.task_id == self.TASK_ID, ) + .execution_options(include_all_attempts=True) ) ) == remaining_ids diff --git a/airflow-core/tests/unit/api_fastapi/execution_api/test_security.py b/airflow-core/tests/unit/api_fastapi/execution_api/test_security.py index 304d1e26f95..e5056e8bd6d 100644 --- a/airflow-core/tests/unit/api_fastapi/execution_api/test_security.py +++ b/airflow-core/tests/unit/api_fastapi/execution_api/test_security.py @@ -339,9 +339,9 @@ class TestAttemptLiveness: if retirement == "retry": assert ( session.scalar( - select(TaskInstance.id).where( - TaskInstance.id == old_id, TaskInstance.working_set.is_(None) - ) + select(TaskInstance.id) + .where(TaskInstance.id == old_id, TaskInstance.working_set.is_(None)) + .execution_options(include_all_attempts=True) ) == old_id ) 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 48bed6901cc..1f1ca7481a6 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 @@ -1451,7 +1451,10 @@ class TestTIUpdateState: assert current.max_tries == expected_max_tries assert current.state is None history = session.scalar( - select(TaskInstance).where(TaskInstance.working_set.is_(None)).where(TaskInstance.id == old_id) + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .where(TaskInstance.id == old_id) + .execution_options(include_all_attempts=True) ) assert history.try_number == 3 assert history.end_date == DEFAULT_END_DATE @@ -1462,7 +1465,10 @@ class TestTIUpdateState: assert (current.id, current.try_number, current.state) == (new_id, 4, None) assert current.max_tries == expected_max_tries assert session.scalars( - select(TaskInstance).where(TaskInstance.working_set.is_(None)).where(TaskInstance.id == old_id) + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .where(TaskInstance.id == old_id) + .execution_options(include_all_attempts=True) ).all() == [history] @pytest.mark.parametrize("first_report", ["api", "executor"]) @@ -1551,6 +1557,7 @@ class TestTIUpdateState: select(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.task_id == "restart_reports") + .execution_options(include_all_attempts=True) ) is None ) @@ -1564,6 +1571,7 @@ class TestTIUpdateState: select(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.task_id == "restart_reports") + .execution_options(include_all_attempts=True) ).one() assert (history.id, history.try_number, history.state) == (old_id, 3, State.FAILED) @@ -2461,6 +2469,7 @@ class TestTIUpdateState: select(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.task_id == ti.task_id, TaskInstance.run_id == ti.run_id) + .execution_options(include_all_attempts=True) ).one() assert tih.id assert tih.id != ti.id @@ -2535,6 +2544,7 @@ class TestTIUpdateState: select(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.task_id == "retired_attempt_report") + .execution_options(include_all_attempts=True) ).all() assert [(attempt.id, attempt.try_number) for attempt in history] == [(old_id, 3)] @@ -2600,6 +2610,7 @@ class TestTIUpdateState: TaskInstance.task_id == ti.task_id, TaskInstance.run_id == ti.run_id, ) + .execution_options(include_all_attempts=True) ).one() assert tih.retry_delay_override == 42.5 assert tih.retry_reason == "Rate limit: backing off" @@ -2652,6 +2663,7 @@ class TestTIUpdateState: TaskInstance.task_id == ti.task_id, TaskInstance.run_id == ti.run_id, ) + .execution_options(include_all_attempts=True) ).one() assert tih.rendered_map_index is None @@ -2690,6 +2702,7 @@ class TestTIUpdateState: TaskInstance.task_id == ti.task_id, TaskInstance.run_id == ti.run_id, ) + .execution_options(include_all_attempts=True) ).one() assert tih.retry_delay_override is None assert tih.retry_reason is None @@ -3435,7 +3448,10 @@ class TestTIHealthEndpoint: assert session.get(TaskInstance, old_ti_id) is not None tih = session.scalar( - select(TaskInstance).where(TaskInstance.working_set.is_(None)).where(TaskInstance.id == old_ti_id) + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .where(TaskInstance.id == old_ti_id) + .execution_options(include_all_attempts=True) ) assert tih is not None @@ -3553,13 +3569,7 @@ class TestTIHealthEndpoint: assert response.status_code == 204 assert len(task_instance_updates) == 1 - assert _where_column_keys(task_instance_updates[0]) == { - "id", - "state", - "hostname", - "pid", - "working_set", - } + assert _where_column_keys(task_instance_updates[0]) == {"id", "state", "hostname", "pid"} assert len(for_update_selects) == 0 session.refresh(ti) assert ti.last_heartbeat_at == new_time diff --git a/airflow-core/tests/unit/jobs/test_scheduler_job.py b/airflow-core/tests/unit/jobs/test_scheduler_job.py index a98decba84a..102aa314599 100644 --- a/airflow-core/tests/unit/jobs/test_scheduler_job.py +++ b/airflow-core/tests/unit/jobs/test_scheduler_job.py @@ -711,6 +711,7 @@ class TestSchedulerJob: select(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.dag_id == dag_id) + .execution_options(include_all_attempts=True) ).one() assert (history.id, history.try_number, history.state) == (retiring_id, 4, State.FAILED) assert history.max_tries == max_tries @@ -722,7 +723,10 @@ class TestSchedulerJob: assert (replacement.id, replacement.try_number, replacement.state) == (replacement_id, 5, None) assert ( session.scalar( - select(func.count()).select_from(TaskInstance).where(TaskInstance.dag_id == dag_id) + select(func.count()) + .select_from(TaskInstance) + .where(TaskInstance.dag_id == dag_id) + .execution_options(include_all_attempts=True) ) == 2 ) @@ -751,6 +755,7 @@ class TestSchedulerJob: .select_from(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.dag_id == ti.dag_id) + .execution_options(include_all_attempts=True) ) == 0 ) @@ -887,8 +892,12 @@ class TestSchedulerJob: self.job_runner.executor.callback_sink.send.assert_not_called() # ti in success state - ti1.state = State.SUCCESS - session.merge(ti1) + session.execute( + update(TaskInstance) + .where(TaskInstance.id == ti1.id) + .values(state=State.SUCCESS) + .execution_options(include_all_attempts=True) + ) session.commit() executor.event_buffer[TaskInstanceUuid(ti1.id)] = State.SUCCESS, None @@ -5844,6 +5853,7 @@ class TestSchedulerJob: TaskInstance.try_number == old_try_number, TaskInstance.id == old_ti_id, ) + .execution_options(include_all_attempts=True) ) is not None ) @@ -5903,7 +5913,10 @@ class TestSchedulerJob: assert (ti.id, ti.state, ti.working_set) == (old_ti_id, State.FAILED, None) tih = session.scalar( - select(TaskInstance).where(TaskInstance.working_set.is_(None)).where(TaskInstance.id == old_ti_id) + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .where(TaskInstance.id == old_ti_id) + .execution_options(include_all_attempts=True) ) assert tih is not None, "TaskInstanceHistory must be created for non-RUNNING retry" assert tih.try_number == 1 @@ -10283,6 +10296,7 @@ class TestSchedulerJob: select(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.dag_id == ti.dag_id) + .execution_options(include_all_attempts=True) ).one() assert (history.id, history.try_number, history.state) == (old_id, 3, State.FAILED) diff --git a/airflow-core/tests/unit/models/test_cleartasks.py b/airflow-core/tests/unit/models/test_cleartasks.py index 91a46ffb218..9a872d3ed9c 100644 --- a/airflow-core/tests/unit/models/test_cleartasks.py +++ b/airflow-core/tests/unit/models/test_cleartasks.py @@ -131,6 +131,7 @@ class TestClearTasks: TaskInstance.run_id == attempt.run_id, TaskInstance.task_id == attempt.task_id, ) + .execution_options(include_all_attempts=True) ) == expected_rows ) @@ -152,7 +153,10 @@ class TestClearTasks: assert (ti.id, ti.try_number, ti.state) == (attempt_id, 4, TaskInstanceState.RESTARTING) assert ( session.scalar( - select(func.count()).select_from(TaskInstance).where(TaskInstance.working_set.is_(None)) + select(func.count()) + .select_from(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .execution_options(include_all_attempts=True) ) == 0 ) @@ -214,8 +218,10 @@ class TestClearTasks: assert (ti.try_number, ti.state) == (2, TaskInstanceState.UP_FOR_RETRY) retry_dep = NotInRetryPeriodDep() assert not retry_dep.is_met(ti, session=session) - history_query = select(TaskInstance.id, TaskInstance.try_number).where( - TaskInstance.dag_id == ti.dag_id, TaskInstance.working_set.is_(None) + history_query = ( + select(TaskInstance.id, TaskInstance.try_number) + .where(TaskInstance.dag_id == ti.dag_id, TaskInstance.working_set.is_(None)) + .execution_options(include_all_attempts=True) ) history_before = session.execute(history_query).mappings().all() assert [(row.id, row.try_number) for row in history_before] == [(failed_id, 1)] @@ -253,6 +259,7 @@ class TestClearTasks: select(func.count()) .select_from(TaskInstance) .where(TaskInstance.dag_id == ti.dag_id, TaskInstance.working_set.is_(None)) + .execution_options(include_all_attempts=True) ) == 0 ) @@ -399,11 +406,27 @@ class TestClearTasks: # but it works for our case because we specifically constructed test DAGS # in the way that those two sort methods are equivalent qry = session.scalars(select(TI).where(TI.dag_id == dag.dag_id).order_by(TI.task_id)).all() - assert session.scalar(select(func.count()).select_from(TI).where(TI.working_set.is_(None))) == 0 + assert ( + session.scalar( + select(func.count()) + .select_from(TI) + .where(TI.working_set.is_(None)) + .execution_options(include_all_attempts=True) + ) + == 0 + ) clear_task_instances(qry, session, dag_run_state=state) session.flush() # 2 TIs were cleared so 2 history records should be created - assert session.scalar(select(func.count()).select_from(TI).where(TI.working_set.is_(None))) == 2 + assert ( + session.scalar( + select(func.count()) + .select_from(TI) + .where(TI.working_set.is_(None)) + .execution_options(include_all_attempts=True) + ) + == 2 + ) session.refresh(dr) @@ -758,7 +781,9 @@ class TestClearTasks: session.flush() session.refresh(dr) - ti_history = session.scalars(select(TI.state).where(TI.working_set.is_(None))).all() + ti_history = session.scalars( + select(TI.state).where(TI.working_set.is_(None)).execution_options(include_all_attempts=True) + ).all() assert ti_history == ([str(state_recorded)] * 2 if state_recorded else []) diff --git a/airflow-core/tests/unit/models/test_dagrun.py b/airflow-core/tests/unit/models/test_dagrun.py index 7edf4d1ac5b..ecf204a16b2 100644 --- a/airflow-core/tests/unit/models/test_dagrun.py +++ b/airflow-core/tests/unit/models/test_dagrun.py @@ -2171,7 +2171,10 @@ def test_restoring_removed_task_allocates_attempt_once(dag_maker, session, try_n session.flush() history = session.scalar( - select(TaskInstance).where(TaskInstance.working_set.is_(None)).where(TaskInstance.id == old_id) + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .where(TaskInstance.id == old_id) + .execution_options(include_all_attempts=True) ) current = dr.get_task_instance("task", map_index=0 if mapped else -1, session=session) assert current.state is None @@ -2214,7 +2217,10 @@ def test_verifying_removed_map_index_does_not_allocate_attempt(dag_maker, sessio assert ti.try_number == 2 assert ( session.scalar( - select(TaskInstance).where(TaskInstance.working_set.is_(None)).where(TaskInstance.id == old_id) + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .where(TaskInstance.id == old_id) + .execution_options(include_all_attempts=True) ) is None ) diff --git a/airflow-core/tests/unit/models/test_task_data.py b/airflow-core/tests/unit/models/test_task_data.py index 8bb7ec34bde..f647ad88585 100644 --- a/airflow-core/tests/unit/models/test_task_data.py +++ b/airflow-core/tests/unit/models/test_task_data.py @@ -197,7 +197,9 @@ def test_joinedload_resolves_relationships_for_rows_from_both_stores( with assert_no_cartesian_products(): rows = session.scalars( - query.options(joinedload(entity.task), joinedload(entity.dag_run).joinedload(DagRun.dag_model)) + query.options( + joinedload(entity.task), joinedload(entity.dag_run).joinedload(DagRun.dag_model) + ).execution_options(include_all_attempts=True) ).unique() by_ti = {row.task_instance_id: row for row in rows} @@ -459,6 +461,7 @@ def test_load_legacy_rendered_fields_fills_only_attempts_without_a_joined_row(ow sa.select(TaskInstance) .where(TaskInstance.id.in_([CURRENT_ID, HISTORY_ID])) .options(joinedload(TaskInstance.rendered_task_instance_fields)) + .execution_options(include_all_attempts=True) ) ) diff --git a/airflow-core/tests/unit/models/test_taskinstance.py b/airflow-core/tests/unit/models/test_taskinstance.py index 016fca7690c..4b56674d5f2 100644 --- a/airflow-core/tests/unit/models/test_taskinstance.py +++ b/airflow-core/tests/unit/models/test_taskinstance.py @@ -653,7 +653,10 @@ class TestTaskInstance: assert ti.try_number == try_number + 1 assert ti.next_retry_datetime() == deadline history = session.scalar( - select(TaskInstance).where(TaskInstance.working_set.is_(None)).where(TaskInstance.id == old_id) + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .where(TaskInstance.id == old_id) + .execution_options(include_all_attempts=True) ) assert history.try_number == try_number @@ -2775,7 +2778,11 @@ class TestTaskInstance: ) == 1 ) - tih = session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(None))).all() + tih = session.scalars( + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .execution_options(include_all_attempts=True) + ).all() assert len(tih) == 1 # the new try_id should be different from what's recorded in tih assert tih[0].id == try_id @@ -2799,7 +2806,11 @@ class TestTaskInstance: ti.retire(reason="retry", session=session) session.flush() - tih = session.scalars(select(TaskInstance).where(TaskInstance.working_set.is_(None))).one() + tih = session.scalars( + select(TaskInstance) + .where(TaskInstance.working_set.is_(None)) + .execution_options(include_all_attempts=True) + ).one() assert tih.state == str(TaskInstanceState.FAILED) assert tih.end_date == archive_time assert tih.duration == (archive_time - start).total_seconds() @@ -2943,20 +2954,34 @@ class TestTaskInstance: assert successor.working_set is True assert successor.state == TaskInstanceState.UP_FOR_RETRY assert successor.end_date == attempt.end_date - assert session.scalar(sa.select(sa.func.count()).select_from(TaskInstance)) == 3 + assert ( + session.scalar( + sa.select(sa.func.count()) + .select_from(TaskInstance) + .execution_options(include_all_attempts=True) + ) + == 3 + ) assert session.scalar(sa.text("SELECT ti_id FROM task_reschedule")) is not None for owner, expected in [(attempt, 1), (successor, 0)]: query = build_xcom_read_query( producer_ids=sa.select(TaskInstance.id).where(TaskInstance.id == owner.id) ) - assert len(session.scalars(query).all()) == expected + assert len(session.scalars(query.execution_options(include_all_attempts=True)).all()) == expected XComModel.set_for_attempt(task_instance_id=attempt.id, key="late", value=1, session=session) assert XComModelV2.get_for_attempt(successor.id, "late", session=session) is None with pytest.raises(ValueError, match="retired"): attempt.prepare_db_for_next_try(session) - assert session.scalar(sa.select(sa.func.count()).select_from(TaskInstance)) == 3 + assert ( + session.scalar( + sa.select(sa.func.count()) + .select_from(TaskInstance) + .execution_options(include_all_attempts=True) + ) + == 3 + ) assert successor.working_set is True def test_retirement_carries_the_note_to_the_successor(self, ownership_session): @@ -3013,10 +3038,70 @@ class TestTaskInstance: session.commit() expected = set() if deleted else {CURRENT_ID, HISTORY_ID} - assert set(session.scalars(select(TaskInstance.id))) == expected + assert ( + set(session.scalars(select(TaskInstance.id).execution_options(include_all_attempts=True))) + == expected + ) assert set(session.scalars(select(XComModelV2.task_instance_id))) == expected assert set(session.scalars(select(RenderedTaskInstanceFields.task_instance_id))) == expected + def test_orm_statements_ignore_retired_attempts_unless_asked(self, ownership_session): + session = ownership_session + session.expunge_all() + include_all_attempts = {"include_all_attempts": True} + + assert set(session.scalars(select(TaskInstance.id))) == {CURRENT_ID} + assert session.scalar(select(func.min(TaskInstance.try_number))) == 2 + assert session.scalar(select(TaskInstance.id).where(TaskInstance.id == HISTORY_ID)) is None + assert set(session.scalars(select(TaskInstance.id).execution_options(**include_all_attempts))) == { + CURRENT_ID, + HISTORY_ID, + } + assert session.scalar( + select(TaskInstance.id) + .where(TaskInstance.id == HISTORY_ID) + .execution_options(**include_all_attempts) + ) + + def test_primary_key_lookups_see_retired_attempts(self, ownership_session): + session = ownership_session + session.expunge_all() + + retired = session.get(TaskInstance, HISTORY_ID) + + assert retired is not None + assert retired.working_set is None + retired.state = TaskInstanceState.SKIPPED + merged = session.merge(retired) + assert merged is retired + + def test_bulk_update_ignores_retired_attempts(self, ownership_session): + session = ownership_session + + updated = session.execute(update(TaskInstance).values(pid=99)).rowcount + + assert updated == 1 + pids = { + row.id: row.pid + for row in session.execute( + select(TaskInstance.id, TaskInstance.pid).execution_options(include_all_attempts=True) + ) + } + assert pids[CURRENT_ID] == 99 + assert pids[HISTORY_ID] != 99 + + def test_retired_attempt_loads_through_relationships_and_refresh(self, ownership_session): + session = ownership_session + retired = session.get(TaskInstance, HISTORY_ID, execution_options={"include_all_attempts": True}) + retired.note = "kept" + session.flush() + session.expire_all() + + note = session.scalar(select(TaskInstanceNote).where(TaskInstanceNote.ti_id == HISTORY_ID)) + assert note.task_instance.id == HISTORY_ID + retired.refresh_from_db(session=session) + assert retired.working_set is None + def test_filter_for_tis_selects_only_the_current_attempt(self, ownership_session): session = ownership_session historical = session.get(TaskInstance, HISTORY_ID) @@ -3050,13 +3135,27 @@ class TestTaskInstance: assert current is attempt assert current.state == TaskInstanceState.RESTARTING assert current.working_set is True - assert session.scalar(sa.select(sa.func.count()).select_from(TaskInstance)) == 2 + assert ( + session.scalar( + sa.select(sa.func.count()) + .select_from(TaskInstance) + .execution_options(include_all_attempts=True) + ) + == 2 + ) successor = current.complete_restart(session=session) assert successor.state is None assert successor.external_executor_id is None assert current.id == CURRENT_ID assert current.working_set is None - assert session.scalar(sa.select(sa.func.count()).select_from(TaskInstance)) == 3 + assert ( + session.scalar( + sa.select(sa.func.count()) + .select_from(TaskInstance) + .execution_options(include_all_attempts=True) + ) + == 3 + ) with pytest.raises(ValueError, match="current restarting"): current.complete_restart(session=session) @@ -3132,7 +3231,11 @@ class TestTaskInstance: if delete_method == "orm": session.delete(target) else: - session.execute(sa.delete(TaskInstance).where(TaskInstance.id == target.id)) + session.execute( + sa.delete(TaskInstance) + .where(TaskInstance.id == target.id) + .execution_options(include_all_attempts=True) + ) session.flush() session.expire_all() @@ -4958,6 +5061,7 @@ def test_failure_listener_receives_failed_try_before_rotation( select(TaskInstance) .where(TaskInstance.working_set.is_(None)) .where(TaskInstance.id == original_id) + .execution_options(include_all_attempts=True) ) assert history.try_number == 1 assert history.state == State.FAILED diff --git a/airflow-core/tests/unit/models/test_trigger.py b/airflow-core/tests/unit/models/test_trigger.py index ab085337e23..97fa53a7133 100644 --- a/airflow-core/tests/unit/models/test_trigger.py +++ b/airflow-core/tests/unit/models/test_trigger.py @@ -510,6 +510,7 @@ def test_submit_event_task_end_failed_respects_retries( TaskInstance.task_id == ti.task_id, TaskInstance.run_id == ti.run_id, ) + .execution_options(include_all_attempts=True) ).all() if expect_history_row: assert len(tih) == 1 diff --git a/airflow-core/tests/unit/utils/test_db_cleanup.py b/airflow-core/tests/unit/utils/test_db_cleanup.py index a8a4bfd91df..713389d7b4c 100644 --- a/airflow-core/tests/unit/utils/test_db_cleanup.py +++ b/airflow-core/tests/unit/utils/test_db_cleanup.py @@ -261,7 +261,9 @@ class TestDBCleanup: def test_task_instance_history_alias_cleans_only_retired_attempts(self, ownership_session): session = ownership_session - session.execute(sa.update(TaskInstance).values(start_date=NOW)) + session.execute( + sa.update(TaskInstance).values(start_date=NOW).execution_options(include_all_attempts=True) + ) session.commit() run_cleanup( @@ -1563,7 +1565,9 @@ class TestDBCleanup: self, ownership_session, parent, dry_run, skip_archive ): session = ownership_session - session.execute(sa.update(TaskInstance).values(start_date=NOW)) + session.execute( + sa.update(TaskInstance).values(start_date=NOW).execution_options(include_all_attempts=True) + ) XComModel.set_for_attempt(task_instance_id=CURRENT_ID, key="new", value=123, session=session) session.execute(sa.update(XComModelV2).values(timestamp=NOW + timedelta(days=30))) session.execute( @@ -1619,7 +1623,9 @@ class TestDBCleanup: @pytest.mark.execution_timeout(10) def test_cleanup_reuses_child_archives_across_parent_batches(self, ownership_session): session = ownership_session - session.execute(sa.update(TaskInstance).values(start_date=NOW)) + session.execute( + sa.update(TaskInstance).values(start_date=NOW).execution_options(include_all_attempts=True) + ) for attempt in session.scalars(sa.select(TaskInstance)): XComModel.set_for_attempt( task_instance_id=attempt.id, key="per_attempt", value=str(attempt.id), session=session
