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..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 @@ -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, @@ -468,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) @@ -529,6 +531,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 +634,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 +666,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 +777,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 +799,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..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 @@ -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,83 @@ 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 + + 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"), [ 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()