Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 15 additions & 3 deletions airflow-core/src/airflow/assets/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,18 @@
log = structlog.get_logger(__name__)


def _sorted_by_dag_id(dags: Collection[DagModel]) -> list[DagModel]:
"""
Order dags deterministically before queueing their AssetDagRunQueue rows.

``dags_to_queue`` is a set, whose iteration order varies between processes. If two
concurrent transactions queue the same dags for one asset in different orders, they take
the per-row locks (from ON CONFLICT / ON DUPLICATE KEY / the SAVEPOINT merge) in opposite
orders and can deadlock. Inserting in a fixed order removes that lock-ordering cycle.
"""
return sorted(dags, key=lambda dag: dag.dag_id)


@contextmanager
def _lock_asset_model(
*,
Expand Down Expand Up @@ -808,7 +820,7 @@ def _queue_dagrun_if_needed(dag: DagModel) -> str | None:
cls.logger().debug("Skipping record %s", item, exc_info=True)
return dag.dag_id

queued_results = (_queue_dagrun_if_needed(dag) for dag in dags_to_queue)
queued_results = (_queue_dagrun_if_needed(dag) for dag in _sorted_by_dag_id(dags_to_queue))
if queued_dag_ids := [r for r in queued_results if r is not None]:
cls.logger().debug("consuming dag ids %s", queued_dag_ids)

Expand All @@ -818,7 +830,7 @@ def _queue_dagruns_nonpartitioned_postgres(
) -> None:
from sqlalchemy.dialects.postgresql import insert

values = [{"target_dag_id": dag.dag_id} for dag in dags_to_queue]
values = [{"target_dag_id": dag.dag_id} for dag in _sorted_by_dag_id(dags_to_queue)]
stmt = insert(AssetDagRunQueue).values(asset_id=asset_id).on_conflict_do_nothing()
session.execute(stmt, values)

Expand All @@ -828,7 +840,7 @@ def _queue_dagruns_nonpartitioned_mysql(
) -> None:
from sqlalchemy.dialects.mysql import insert

values = [{"target_dag_id": dag.dag_id} for dag in dags_to_queue]
values = [{"target_dag_id": dag.dag_id} for dag in _sorted_by_dag_id(dags_to_queue)]
stmt = insert(AssetDagRunQueue).values(asset_id=asset_id)
# MySQL has no "ON CONFLICT DO NOTHING"; a no-op ON DUPLICATE KEY UPDATE turns a
# conflicting (asset_id, target_dag_id) row into a no-op rather than an error,
Expand Down
39 changes: 38 additions & 1 deletion airflow-core/tests/unit/assets/test_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@
from airflow import settings
from airflow._shared.observability.metrics.base_stats_logger import StatsLogger
from airflow._shared.timezones import timezone
from airflow.assets.manager import AssetManager
from airflow.assets.manager import AssetManager, _sorted_by_dag_id
from airflow.models.asset import (
AssetAliasModel,
AssetDagRunQueue,
Expand Down Expand Up @@ -252,6 +252,43 @@ def test_queue_dagruns_nonpartitioned_mysql_builds_upsert(self):
assert "ON DUPLICATE KEY UPDATE" in compiled
assert values == [{"target_dag_id": "dag1"}]

def test_sorted_by_dag_id_orders_deterministically(self):
dags = [DagModel(dag_id="dag_c"), DagModel(dag_id="dag_a"), DagModel(dag_id="dag_b")]
assert [d.dag_id for d in _sorted_by_dag_id(dags)] == ["dag_a", "dag_b", "dag_c"]

@pytest.mark.parametrize(
"helper", ["_queue_dagruns_nonpartitioned_postgres", "_queue_dagruns_nonpartitioned_mysql"]
)
def test_queue_dagruns_single_statement_inserts_in_helper_order(self, helper):
"""Both single-statement paths must build the insert from _sorted_by_dag_id, not set order."""
dags = {DagModel(dag_id="dag_a"), DagModel(dag_id="dag_b"), DagModel(dag_id="dag_c")}
ordered = [DagModel(dag_id="dag_c"), DagModel(dag_id="dag_b"), DagModel(dag_id="dag_a")]
session = mock.MagicMock(spec=Session)

with mock.patch("airflow.assets.manager._sorted_by_dag_id", return_value=ordered) as sorter:
getattr(AssetManager, helper)(asset_id=1, dags_to_queue=dags, session=session)

sorter.assert_called_once_with(dags)
# The insert order must follow the helper's return, not the set's iteration order.
_, values = session.execute.call_args.args
assert values == [{"target_dag_id": d.dag_id} for d in ordered]

def test_queue_dagruns_slow_path_merges_in_helper_order(self):
"""The per-row SAVEPOINT path must merge from _sorted_by_dag_id, not set order."""
dags = {DagModel(dag_id="dag_a"), DagModel(dag_id="dag_b"), DagModel(dag_id="dag_c")}
ordered = [DagModel(dag_id="dag_c"), DagModel(dag_id="dag_b"), DagModel(dag_id="dag_a")]
session = mock.MagicMock(spec=Session)

with mock.patch("airflow.assets.manager._sorted_by_dag_id", return_value=ordered) as sorter:
AssetManager._queue_dagruns_nonpartitioned_slow_path(
asset_id=1, dags_to_queue=dags, session=session
)

sorter.assert_called_once_with(dags)
# The merge order must follow the helper's return, not the set's iteration order.
merged = [call.args[0].target_dag_id for call in session.merge.call_args_list]
assert merged == [d.dag_id for d in ordered]

def test_register_asset_change_notifies_asset_listener(
self, session, mock_task_instance, testing_dag_bundle, listener_manager
):
Expand Down
Loading