From 88a46bec167e7ec99fd2be74ff7caabc8ecb9a2c Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Tue, 4 Aug 2026 14:55:12 +0800 Subject: [PATCH 1/2] [v3-3-test] Release TI lock before asset listener callbacks (#70951) * Release TI lock before asset listener callbacks Asset registration on the task-success path (ti_update_state) ran the listener hooks synchronously inside the transaction holding a row lock on the task_instance table. A slow listener, multiplied across a large fan-out of asset events, could hold that lock for minutes, causing statement timeouts. The listener hooks are now deferred until the end of the endpoint instead of executed inline during asset event creation. Registration writes to the database still happen under the caller's transaction, so durability is unchanged; this only moves the best-effort listener hooks off the lock. * Optimize asset alias assoc insert * Fix exhausted iterator reuse bug * Test asset reg callback cases (cherry picked from commit 79db99500064aa801bcfad1b6f91b8be867f822c) Co-authored-by: Tzu-ping Chung --- .../execution_api/routes/task_instances.py | 20 ++++-- airflow-core/src/airflow/assets/manager.py | 63 ++++++++++++------- .../src/airflow/models/taskinstance.py | 13 +++- .../versions/head/test_task_instances.py | 36 +++++++++++ .../tests/unit/assets/test_manager.py | 38 +++++++++++ 5 files changed, 138 insertions(+), 32 deletions(-) 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 34c3dc35406f8..b45aa4e265560 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 @@ -21,7 +21,7 @@ import itertools import json from collections import defaultdict -from collections.abc import Iterator +from collections.abc import Callable, Iterator, Sequence from typing import TYPE_CHECKING, Annotated, Any, NoReturn, cast from uuid import UUID @@ -450,8 +450,9 @@ def ti_update_state( data["_rendered_map_index"] = data.pop("rendered_map_index") query = update(TI).where(TI.id == task_instance_id).values(data) + asset_callbacks: Sequence[Callable[[], None]] = () try: - query, updated_state = _create_ti_state_update_query_and_update_state( + query, updated_state, asset_callbacks = _create_ti_state_update_query_and_update_state( ti_patch_payload=ti_patch_payload, task_instance_id=task_instance_id, session=session, @@ -529,6 +530,12 @@ def ti_update_state( task_id=task_id, ) + # Release the task_instance row lock before running listener callbacks. + session.commit() + + for callback in asset_callbacks: + callback() + def _emit_task_span(ti, state): # just to be safe @@ -626,7 +633,8 @@ def _create_ti_state_update_query_and_update_state( session: SessionDep, dag_bag: DagBagDep, dag_id: str, -) -> tuple[Update, TaskInstanceState]: +) -> tuple[Update, TaskInstanceState, Sequence[Callable[[], None]]]: + asset_callbacks: Sequence[Callable[[], None]] = () if isinstance(ti_patch_payload, (TITerminalStatePayload, TIRetryStatePayload, TISuccessStatePayload)): ti = session.get(TI, task_instance_id, with_for_update={"of": TI}) updated_state = TaskInstanceState(ti_patch_payload.state.value) @@ -657,7 +665,7 @@ def _create_ti_state_update_query_and_update_state( query = query.values(retry_delay_override=retry_delay_override, retry_reason=retry_reason) elif isinstance(ti_patch_payload, TISuccessStatePayload): if ti is not None: - TI.register_asset_changes_in_db( + asset_callbacks = TI.register_asset_changes_in_db( ti, ti_patch_payload.task_outlets, ti_patch_payload.outlet_events, @@ -768,7 +776,7 @@ def _create_ti_state_update_query_and_update_state( ti = session.get(TI, task_instance_id, with_for_update={"of": TI}) if ti is not None: _handle_fail_fast_for_dag(ti=ti, dag_id=dag_id, session=session, dag_bag=dag_bag) - return query, TaskInstanceState.FAILED + return query, TaskInstanceState.FAILED, () actual_start_date = timezone.utcnow() session.add( @@ -790,7 +798,7 @@ def _create_ti_state_update_query_and_update_state( else: raise ValueError(f"Unexpected Payload Type {type(ti_patch_payload)}") - return query, updated_state + return query, updated_state, asset_callbacks @ti_id_router.patch( diff --git a/airflow-core/src/airflow/assets/manager.py b/airflow-core/src/airflow/assets/manager.py index c3491574e9148..0c7fa50d5af97 100644 --- a/airflow-core/src/airflow/assets/manager.py +++ b/airflow-core/src/airflow/assets/manager.py @@ -17,12 +17,13 @@ # under the License. from __future__ import annotations -from collections.abc import Collection, Iterable +from collections.abc import Callable, Collection, Iterable from contextlib import contextmanager +from functools import partial from typing import TYPE_CHECKING import structlog -from sqlalchemy import exc, or_, select +from sqlalchemy import exc, insert, or_, select from sqlalchemy.orm import joinedload from airflow._shared.observability.metrics import stats @@ -41,6 +42,7 @@ DagScheduleAssetUriReference, PartitionedAssetKeyLog, TaskOutletAssetReference, + asset_alias_asset_event_association_table, ) from airflow.models.log import Log from airflow.timetables.base import compute_rollup_fingerprint @@ -281,6 +283,7 @@ def register_asset_change( api_user_teams: set[str] | None = None, api_allow_consumer_teams: list[str] | None = None, api_allow_global_consumers: bool = True, + callback_sink: list[Callable[[], None]] | None = None, **kwargs, ) -> AssetEvent | None: """ @@ -306,6 +309,8 @@ def register_asset_change( Only used when source_is_api=True. :param api_allow_global_consumers: Whether teamless consumers are allowed for an API-triggered event. Only used when source_is_api=True. Defaults to True. + :param callback_sink: If specified, registration callbacks are added + into the list instead of executed inline. """ from airflow.models.dag import DagModel @@ -350,17 +355,26 @@ def register_asset_change( dags_to_queue_from_asset_alias = set() if source_alias_names: - asset_alias_models: Iterable[AssetAliasModel] = session.scalars( - select(AssetAliasModel) - .where(AssetAliasModel.name.in_(source_alias_names)) - .options( - joinedload(AssetAliasModel.scheduled_dags).joinedload(DagScheduleAssetAliasReference.dag) + asset_alias_models = ( + session.scalars( + select(AssetAliasModel) + .where(AssetAliasModel.name.in_(source_alias_names)) + .options( + joinedload(AssetAliasModel.scheduled_dags).joinedload( + DagScheduleAssetAliasReference.dag + ) + ) ) - ).unique() + .unique() + .all() + ) for asset_alias_model in asset_alias_models: - asset_alias_model.asset_events.append(asset_event) - session.add(asset_alias_model) + session.execute( + insert(asset_alias_asset_event_association_table).values( + alias_id=asset_alias_model.id, event_id=asset_event.id + ) + ) dags_to_queue_from_asset_alias |= { alias_ref.dag @@ -386,20 +400,23 @@ def register_asset_change( ) asset = asset_model.to_serialized() - cls.notify_asset_changed(asset=asset) - cls.nofity_asset_event_emitted( - asset_event=ListenerAssetEvent( - asset=asset, - extra=asset_event.extra, - source_dag_id=asset_event.source_dag_id, - source_task_id=asset_event.source_task_id, - source_run_id=asset_event.source_run_id, - source_map_index=asset_event.source_map_index, - source_aliases=[aam.to_serialized() for aam in asset_alias_models], - partition_key=partition_key, - partition_date=partition_date, - ) + listener_asset_event = ListenerAssetEvent( + asset=asset, + extra=asset_event.extra, + source_dag_id=asset_event.source_dag_id, + source_task_id=asset_event.source_task_id, + source_run_id=asset_event.source_run_id, + source_map_index=asset_event.source_map_index, + source_aliases=[aam.to_serialized() for aam in asset_alias_models], + partition_key=partition_key, + partition_date=partition_date, ) + if callback_sink is None: + cls.notify_asset_changed(asset=asset) + cls.nofity_asset_event_emitted(asset_event=listener_asset_event) + else: + callback_sink.append(partial(cls.notify_asset_changed, asset=asset)) + callback_sink.append(partial(cls.nofity_asset_event_emitted, asset_event=listener_asset_event)) team_name = None if task_instance and conf.getboolean("core", "multi_team"): diff --git a/airflow-core/src/airflow/models/taskinstance.py b/airflow-core/src/airflow/models/taskinstance.py index f2bd3f18bebaf..7dae21baa2ccc 100644 --- a/airflow-core/src/airflow/models/taskinstance.py +++ b/airflow-core/src/airflow/models/taskinstance.py @@ -24,7 +24,7 @@ import math import warnings from collections import defaultdict -from collections.abc import Collection, Iterable +from collections.abc import Callable, Collection, Iterable, Sequence from datetime import datetime, timedelta from typing import TYPE_CHECKING, Any, NamedTuple from urllib.parse import quote @@ -1524,14 +1524,14 @@ def register_asset_changes_in_db( outlet_events: list[dict[str, Any]], *, session: Session = NEW_SESSION, - ) -> None: + ) -> Sequence[Callable[[], None]]: # Fast path: a task with no outlets and no outlet events has nothing to # register. Returning early avoids the AssetModel lookup below (which # would run with empty IN () clauses) and all downstream work. This is # the common case -- most tasks declare no outlets -- and it sits on the # task-success path that gates scheduling the next task. if not task_outlets and not outlet_events: - return + return () from airflow.serialization.definitions.assets import ( SerializedAsset, @@ -1557,6 +1557,7 @@ def register_asset_changes_in_db( dag_run_partition_key = ti.dag_run.partition_key dag_run_partition_date = ti.dag_run.partition_date + callback_sink: list[Callable[[], None]] = [] asset_keys = { SerializedAssetUniqueKey(o.name, o.uri) for o in task_outlets @@ -1592,6 +1593,7 @@ def _register(am: AssetModel, key: SerializedAssetUniqueKey) -> None: extra=None, partition_key=dag_run_partition_key, partition_date=dag_run_partition_date, + callback_sink=callback_sink, session=session, ) return @@ -1619,6 +1621,7 @@ def _register(am: AssetModel, key: SerializedAssetUniqueKey) -> None: extra=payload.extra, partition_key=effective_pk, partition_date=payload_partition_date, + callback_sink=callback_sink, session=session, ) @@ -1703,6 +1706,7 @@ def _asset_event_extras_from_aliases() -> dict[tuple[SerializedAssetUniqueKey, s extra=asset_event_extra, partition_key=dag_run_partition_key, partition_date=dag_run_partition_date, + callback_sink=callback_sink, session=session, ) if event is None: @@ -1716,9 +1720,12 @@ def _asset_event_extras_from_aliases() -> dict[tuple[SerializedAssetUniqueKey, s extra=asset_event_extra, partition_key=dag_run_partition_key, partition_date=dag_run_partition_date, + callback_sink=callback_sink, session=session, ) + return callback_sink + @provide_session def update_rtif(self, rendered_fields, *, session: Session = NEW_SESSION): from airflow.models.renderedtifields import RenderedTaskInstanceFields 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 8a152bebe0d3f..a744db8bf10fe 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 @@ -66,6 +66,7 @@ clear_db_serialized_dags, clear_rendered_ti_fields, ) +from unit.listeners import asset_listener if TYPE_CHECKING: from airflow.sdk.api.client import Client @@ -1295,6 +1296,41 @@ def test_ti_update_state_to_success_with_asset_events( assert event[0].asset == AssetModel(name="my-task", uri="s3://bucket/my-task", extra={}) assert event[0].extra == expected_extra + def test_ti_update_state_to_success_runs_deferred_asset_listener_callbacks( + self, client, session, create_task_instance, listener_manager + ): + """The success endpoint runs the deferred asset listener callbacks after committing.""" + asset_listener.clear() + listener_manager(asset_listener) + + asset = AssetModel(id=1, name="my-task", uri="s3://bucket/my-task", group="asset", extra={}) + session.add_all([asset, AssetActive.for_asset(asset)]) + + ti = create_task_instance( + task_id="test_ti_update_state_to_success_runs_deferred_asset_listener_callbacks", + start_date=DEFAULT_START_DATE, + state=State.RUNNING, + ) + session.commit() + + response = client.patch( + f"/execution/task-instances/{ti.id}/state", + json={ + "state": "success", + "end_date": DEFAULT_END_DATE.isoformat(), + "task_outlets": [{"name": "my-task", "uri": "s3://bucket/my-task", "type": "Asset"}], + "outlet_events": [], + }, + ) + + assert response.status_code == 204 + + # Notifications are deferred during registration and run by the endpoint after the + # TI state is committed (and the task_instance row lock released). + assert len(asset_listener.changed) == 1 + assert asset_listener.changed[0].uri == "s3://bucket/my-task" + assert len(asset_listener.emitted) == 1 + @pytest.mark.parametrize( ("outlet_events", "expected_extra"), [ diff --git a/airflow-core/tests/unit/assets/test_manager.py b/airflow-core/tests/unit/assets/test_manager.py index b788b9ab28699..2b876e22b132f 100644 --- a/airflow-core/tests/unit/assets/test_manager.py +++ b/airflow-core/tests/unit/assets/test_manager.py @@ -236,6 +236,44 @@ def test_register_asset_change_notifies_asset_listener( assert len(asset_listener.changed) == 1 assert asset_listener.changed[0].uri == asset.uri + def test_register_asset_change_defers_notifications_to_callback_sink( + self, session, mock_task_instance, testing_dag_bundle, listener_manager + ): + asset_manager = AssetManager() + asset_listener.clear() + listener_manager(asset_listener) + + bundle_name = "testing" + + asset = Asset(uri="test://asset1", name="test_asset_1") + dag1 = DagModel(dag_id="dag3", bundle_name=bundle_name) + session.add(dag1) + + asm = AssetModel(uri="test://asset1/", name="test_asset_1", group="asset") + session.add(asm) + asm.scheduled_dags = [DagScheduleAssetReference(dag_id=dag1.dag_id)] + session.flush() + + # When a callback_sink is supplied, listener notifications are collected into it + # instead of firing inline, so the caller can run them after releasing the lock. + callback_sink: list = [] + asset_manager.register_asset_change( + task_instance=mock_task_instance, + asset=asset, + session=session, + callback_sink=callback_sink, + ) + session.flush() + + assert asset_listener.changed == [] + assert callback_sink + + # Running the collected callbacks fires the listeners. + for callback in callback_sink: + callback() + assert len(asset_listener.changed) == 1 + assert asset_listener.changed[0].uri == asset.uri + def test_create_assets_notifies_asset_listener(self, session, listener_manager): asset_manager = AssetManager() asset_listener.clear() From b5e4f91bc1ac252ba14c1461dd28a41c0235de90 Mon Sep 17 00:00:00 2001 From: Tzu-ping Chung Date: Tue, 4 Aug 2026 19:17:13 +0800 Subject: [PATCH 2/2] [v3-3-test] Explicitly rollback on task state update exception (#71076) --- .../execution_api/routes/task_instances.py | 1 + .../versions/head/test_task_instances.py | 42 +++++++++++++++++++ 2 files changed, 43 insertions(+) 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 b45aa4e265560..41ecf49b053fb 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 @@ -469,6 +469,7 @@ def ti_update_state( "Error updating Task Instance state. Setting the task to failed.", payload=ti_patch_payload, ) + session.rollback() ti = session.get(TI, task_instance_id, with_for_update={"of": TI}) if session.bind is not None: query = TI.duration_expression_update(timezone.utcnow(), query, session.bind) 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 a744db8bf10fe..bb3c0f7e5a785 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 @@ -1331,6 +1331,48 @@ def test_ti_update_state_to_success_runs_deferred_asset_listener_callbacks( assert asset_listener.changed[0].uri == "s3://bucket/my-task" assert len(asset_listener.emitted) == 1 + def test_ti_update_state_to_success_rolls_back_partial_asset_registration( + self, client, session, create_task_instance + ): + """A failure partway through asset registration rolls back the partial writes. + + The endpoint's exception handler marks the TI failed and commits; the explicit rollback + ensures an asset event flushed before the failure is not committed alongside it. + """ + asset = AssetModel(id=1, name="my-task", uri="s3://bucket/my-task", group="asset", extra={}) + session.add_all([asset, AssetActive.for_asset(asset)]) + + ti = create_task_instance( + task_id="test_ti_update_state_to_success_rolls_back_partial_asset_registration", + start_date=DEFAULT_START_DATE, + state=State.RUNNING, + ) + session.commit() + + def _partial_then_fail(ti, task_outlets, outlet_events, *, session): + # Simulate a half-way registration: an asset event is flushed, then registration + # fails before the endpoint commits. + session.add(AssetEvent(asset_id=asset.id)) + session.flush() + raise RuntimeError("boom partway through outlets") + + with mock.patch.object(TaskInstance, "register_asset_changes_in_db", side_effect=_partial_then_fail): + response = client.patch( + f"/execution/task-instances/{ti.id}/state", + json={ + "state": "success", + "end_date": DEFAULT_END_DATE.isoformat(), + "task_outlets": [{"name": "my-task", "uri": "s3://bucket/my-task", "type": "Asset"}], + "outlet_events": [], + }, + ) + + assert response.status_code == 204 + session.expire_all() + # The partially-written asset event was rolled back, and the TI is marked failed. + assert session.scalars(select(AssetEvent)).all() == [] + assert session.get(TaskInstance, ti.id).state == State.FAILED + @pytest.mark.parametrize( ("outlet_events", "expected_extra"), [