From 7efd83e1fbe3576670a7983e3fd79a6f705af758 Mon Sep 17 00:00:00 2001 From: Anmol Ahuja Date: Thu, 6 Aug 2026 13:57:43 +0530 Subject: [PATCH] Pass precomputed hash to SerializedDagModel constructor write_dag() computes cls.hash(dag.data) to compare against the stored hash. When the DAG has changed, it then calls cls(dag) which recomputes the same hash internally. Pass the already-known hash via _precomputed_hash to skip the redundant computation. When reused_deadline_data mutates dag.data after the initial hash, the precomputed hash is stale, so the hash is recomputed once more right after the mutation before being passed in. Addresses the feedback from #70715. --- .../src/airflow/models/serialized_dag.py | 13 +- .../tests/unit/models/test_serialized_dag.py | 117 ++++++++++++++++++ 2 files changed, 126 insertions(+), 4 deletions(-) diff --git a/airflow-core/src/airflow/models/serialized_dag.py b/airflow-core/src/airflow/models/serialized_dag.py index f3f5b78309680..81ebdf799f1ff 100644 --- a/airflow-core/src/airflow/models/serialized_dag.py +++ b/airflow-core/src/airflow/models/serialized_dag.py @@ -345,10 +345,13 @@ class SerializedDagModel(Base): load_op_links = True - def __init__(self, dag: LazyDeserializedDAG) -> None: + def __init__(self, dag: LazyDeserializedDAG, *, _precomputed_hash: str | None = None) -> None: self.dag_id = dag.dag_id dag_data = dag.data - self.dag_hash = SerializedDagModel.hash(dag_data) + if _precomputed_hash is not None: + self.dag_hash = _precomputed_hash + else: + self.dag_hash = SerializedDagModel.hash(dag_data) # partially ordered json data dag_data_json = json.dumps(dag_data, sort_keys=True).encode("utf-8") @@ -727,7 +730,7 @@ def write_dag( # This is for dynamic DAGs that the hashes changes often. We should update # the serialized dag, the dag_version and the dag_code instead of a new version # if the dag_version is not associated with any task instances - new_serialized_dag = cls(dag) + new_serialized_dag = cls(dag, _precomputed_hash=new_dag_hash) # Use direct UPDATE to avoid loading the full serialized DAG result = session.execute( @@ -778,8 +781,10 @@ def write_dag( if reused_deadline_data: deadline_uuid_mapping = {str(uuid6.uuid7()): data for data in reused_deadline_data.values()} dag.data["dag"]["deadline"] = list(deadline_uuid_mapping.keys()) + # The deadline UUIDs changed after new_dag_hash was computed, so refresh it. + new_dag_hash = cls.hash(dag.data) - new_serialized_dag = cls(dag) + new_serialized_dag = cls(dag, _precomputed_hash=new_dag_hash) new_serialized_dag.dag_version = dagv session.add(new_serialized_dag) diff --git a/airflow-core/tests/unit/models/test_serialized_dag.py b/airflow-core/tests/unit/models/test_serialized_dag.py index fd3568f34619c..5d9e1b40b1af8 100644 --- a/airflow-core/tests/unit/models/test_serialized_dag.py +++ b/airflow-core/tests/unit/models/test_serialized_dag.py @@ -19,6 +19,7 @@ from __future__ import annotations +import copy import logging from datetime import timedelta from unittest import mock @@ -1234,3 +1235,119 @@ def test_deadline_reuse_skips_write_when_hash_matches(self, testing_dag_bundle, alert = session.scalar(select(DAM).where(DAM.serialized_dag_id == orig_serdag.id)) assert alert is not None assert alert.id == orig_alert.id + + +class TestPrecomputedHash: + """Tests for the ``_precomputed_hash`` parameter of ``SerializedDagModel.__init__``.""" + + DAG_DATA = {"dag": {"dag_id": "precomputed_dag", "fileloc": "/tmp/precomputed_dag.py", "tasks": []}} + + @pytest.fixture(autouse=True) + def setup_test_cases(self): + db.clear_db_dags() + db.clear_db_runs() + db.clear_db_serialized_dags() + yield + db.clear_db_serialized_dags() + + def _make_lazy_dag(self): + lazy_dag = mock.MagicMock(spec=LazyDeserializedDAG) + lazy_dag.dag_id = "precomputed_dag" + lazy_dag.data = copy.deepcopy(self.DAG_DATA) + return lazy_dag + + @mock.patch.object(SDM, "hash", autospec=True) + def test_precomputed_hash_skips_hash_call(self, mock_hash): + lazy_dag = self._make_lazy_dag() + model = SDM(lazy_dag, _precomputed_hash="abc123") + mock_hash.assert_not_called() + assert model.dag_hash == "abc123" + + def test_precomputed_hash_produces_identical_storage(self): + real_hash = SDM.hash(copy.deepcopy(self.DAG_DATA)) + with_hash = SDM(self._make_lazy_dag(), _precomputed_hash=real_hash) + without_hash = SDM(self._make_lazy_dag()) + assert with_hash.dag_hash == without_hash.dag_hash == real_hash + assert with_hash._data == without_hash._data + assert with_hash._data_compressed == without_hash._data_compressed + + @mock.patch.object(SDM, "hash", autospec=True, return_value="computed") + def test_without_precomputed_hash_calls_hash_normally(self, mock_hash): + lazy_dag = self._make_lazy_dag() + model = SDM(lazy_dag) + mock_hash.assert_called_once_with(lazy_dag.data) + assert model.dag_hash == "computed" + + def test_write_dag_changed_dag_passes_precomputed_hash(self, dag_maker, session): + with dag_maker(dag_id="test_precomputed_hash_write", session=session) as dag: + EmptyOperator(task_id="task1") + SDM.write_dag(LazyDeserializedDAG.from_dag(dag), bundle_name="dag_maker", session=session) + session.commit() + + # Add a task so the hash changes and write_dag reaches a cls(dag) call site. + EmptyOperator(task_id="task2", dag=dag) + lazy_dag = LazyDeserializedDAG.from_dag(dag) + expected_hash = SDM.hash(lazy_dag.data) + + orig_init = SDM.__init__ + with mock.patch.object(SDM, "__init__", autospec=True, side_effect=orig_init) as mock_init: + assert SDM.write_dag(lazy_dag, bundle_name="dag_maker", session=session) is True + session.commit() + + assert mock_init.call_args.kwargs["_precomputed_hash"] == expected_hash + assert SDM.get(dag.dag_id, session=session).dag_hash == expected_hash + + def test_deadline_mutation_does_not_pass_stale_hash(self, testing_dag_bundle, session): + """ + When deadline UUIDs are regenerated after the hash comparison, the stored + hash must reflect the regenerated UUIDs, not the stale precomputed hash. + """ + dag_id = "test_precomputed_hash_deadline" + + dag = DAG( + dag_id=dag_id, + deadline=DeadlineAlert( + reference=DeadlineReference.DAGRUN_QUEUED_AT, + interval=timedelta(minutes=5), + callback=AsyncCallback(empty_callback_for_deadline), + ), + ) + EmptyOperator(task_id="task1", dag=dag) + scheduler_dag = sync_dag_to_db(dag, session=session) + + # Create a dagrun so the existing dag_version has task instances, + # forcing write_dag into the INSERT branch where reused deadline + # UUIDs are regenerated after new_dag_hash was computed. + scheduler_dag.create_dagrun( + run_id="test1", + run_after=DEFAULT_DATE, + state=DagRunState.QUEUED, + logical_date=DEFAULT_DATE, + data_interval=(DEFAULT_DATE, DEFAULT_DATE), + triggered_by=DagRunTriggeredByType.TEST, + run_type=DagRunType.MANUAL, + ) + session.commit() + + orig_serdag = session.scalar(select(SDM).where(SDM.dag_id == dag_id).order_by(SDM.created_at.desc())) + orig_deadline_uuids = orig_serdag.data["dag"]["deadline"] + + # Add a second task (non-deadline change) — deadline definitions are + # unchanged, so write_dag reuses them and then regenerates fresh UUIDs. + EmptyOperator(task_id="task2", dag=dag) + SDM.write_dag(LazyDeserializedDAG.from_dag(dag), bundle_name="testing", session=session) + session.commit() + + new_serdag = session.scalar(select(SDM).where(SDM.dag_id == dag_id).order_by(SDM.created_at.desc())) + new_deadline_uuids = new_serdag.data["dag"]["deadline"] + assert new_serdag.id != orig_serdag.id + assert new_deadline_uuids != orig_deadline_uuids + + # The stored hash must match the stored data (with the fresh UUIDs). + assert new_serdag.dag_hash == SDM.hash(new_serdag.data) + + # And must NOT match the stale hash computed before the UUID regeneration, + # i.e. the hash of the same data but with the old (reused) UUIDs. + stale_data = copy.deepcopy(new_serdag.data) + stale_data["dag"]["deadline"] = orig_deadline_uuids + assert new_serdag.dag_hash != SDM.hash(stale_data)