diff --git a/airflow-core/src/airflow/models/taskinstance.py b/airflow-core/src/airflow/models/taskinstance.py index ff610f2146588..a654e59a23f29 100644 --- a/airflow-core/src/airflow/models/taskinstance.py +++ b/airflow-core/src/airflow/models/taskinstance.py @@ -248,10 +248,14 @@ def _recalculate_dagrun_queued_at_deadlines( if not results: return + # Local import to avoid a circular import between models and serialization. + from airflow.serialization.decoders import decode_deadline_alert_model, resolve_deadline_alert_interval + for deadline, deadline_alert in results: - # We can't use evaluate_with() since the new queued_at is not written to the DB yet. - deadline_interval = timedelta(seconds=deadline_alert.interval) - new_deadline_time = new_queued_at + deadline_interval + # We can't use evaluate_with() since the new queued_at is not written to the DB yet, and + # interval is stored as JSON, so it has to be decoded rather than passed to timedelta(). + interval = resolve_deadline_alert_interval(decode_deadline_alert_model(deadline_alert)) + new_deadline_time = new_queued_at + interval log.debug( "Recalculating deadline %s for DagRun %s.%s: old=%s, new=%s", diff --git a/airflow-core/src/airflow/serialization/decoders.py b/airflow-core/src/airflow/serialization/decoders.py index 49f88480b57cb..c6d7a2cd81612 100644 --- a/airflow-core/src/airflow/serialization/decoders.py +++ b/airflow-core/src/airflow/serialization/decoders.py @@ -24,6 +24,7 @@ import dateutil.relativedelta from airflow._shared.module_loading import import_string +from airflow.sdk.definitions.deadline import VariableInterval from airflow.serialization.definitions.assets import ( SerializedAsset, SerializedAssetAlias, @@ -52,6 +53,7 @@ ) if TYPE_CHECKING: + from airflow.models.deadline_alert import DeadlineAlert as DeadlineAlertModel from airflow.partition_mappers.base import PartitionMapper from airflow.partition_mappers.wait_policy import WaitPolicy from airflow.partition_mappers.window import Window @@ -181,13 +183,12 @@ def decode_deadline_reference(reference_data: dict): return reference_class.deserialize_reference(reference_data) -def decode_deadline_alert(encoded_data: dict): +def decode_deadline_alert(encoded_data: dict) -> SerializedDeadlineAlert: """ Decode a previously serialized deadline alert. :meta private: """ - from airflow.sdk.definitions.deadline import VariableInterval from airflow.sdk.serde import deserialize data = encoded_data.get(Encoding.VAR, encoded_data) @@ -224,6 +225,35 @@ def decode_deadline_alert(encoded_data: dict): ) +def decode_deadline_alert_model(deadline_alert: DeadlineAlertModel) -> SerializedDeadlineAlert: + """ + Decode a ``DeadlineAlert`` ORM row into its serialized representation. + + :meta private: + """ + return decode_deadline_alert( + { + DeadlineAlertFields.REFERENCE: deadline_alert.reference, + DeadlineAlertFields.INTERVAL: deadline_alert.interval, + DeadlineAlertFields.CALLBACK: deadline_alert.callback_def, + } + ) + + +def resolve_deadline_alert_interval(alert: SerializedDeadlineAlert) -> datetime.timedelta: + """ + Resolve a decoded alert's interval to a ``timedelta``. + + A ``VariableInterval`` reads its Airflow Variable here, so this is only called at the point + a deadline is actually calculated. + + :meta private: + """ + if isinstance(alert.interval, VariableInterval): + return alert.interval.resolve() + return alert.interval + + def decode_timetable(var: dict[str, Any]) -> CoreTimetable: """ Decode a previously serialized timetable. diff --git a/airflow-core/src/airflow/serialization/definitions/dag.py b/airflow-core/src/airflow/serialization/definitions/dag.py index 8ed0fee2ccabd..db63d7326bdc6 100644 --- a/airflow-core/src/airflow/serialization/definitions/dag.py +++ b/airflow-core/src/airflow/serialization/definitions/dag.py @@ -50,11 +50,9 @@ from airflow.models.deadline_alert import DeadlineAlert as DeadlineAlertModel from airflow.models.taskinstancekey import TaskInstanceKey from airflow.models.tasklog import LogTemplate -from airflow.sdk.definitions.deadline import VariableInterval -from airflow.serialization.decoders import decode_deadline_alert -from airflow.serialization.definitions.deadline import DeadlineAlertFields, SerializedReferenceModels +from airflow.serialization.decoders import decode_deadline_alert_model, resolve_deadline_alert_interval +from airflow.serialization.definitions.deadline import SerializedReferenceModels from airflow.serialization.definitions.param import SerializedParamsDict -from airflow.serialization.enums import DagAttributeTypes as DAT, Encoding from airflow.timetables.base import DagRunInfo, DataInterval, TimeRestriction from airflow.utils.helpers import prune_dict from airflow.utils.session import NEW_SESSION, provide_session @@ -741,21 +739,8 @@ def _process_dagrun_deadline_alerts( if not deadline_alert: continue - deserialized_deadline_alert = decode_deadline_alert( - { - Encoding.TYPE: DAT.DEADLINE_ALERT, - Encoding.VAR: { - DeadlineAlertFields.REFERENCE: deadline_alert.reference, - DeadlineAlertFields.INTERVAL: deadline_alert.interval, - DeadlineAlertFields.CALLBACK: deadline_alert.callback_def, - }, - } - ) - - interval = deserialized_deadline_alert.interval - - if isinstance(interval, VariableInterval): - interval = interval.resolve() + deserialized_deadline_alert = decode_deadline_alert_model(deadline_alert) + interval = resolve_deadline_alert_interval(deserialized_deadline_alert) if isinstance(deserialized_deadline_alert.reference, SerializedReferenceModels.TYPES.DAGRUN): deadline_time = deserialized_deadline_alert.reference.evaluate_with( diff --git a/airflow-core/tests/unit/models/test_taskinstance.py b/airflow-core/tests/unit/models/test_taskinstance.py index 6a348a3f2ca26..81d3390f73955 100644 --- a/airflow-core/tests/unit/models/test_taskinstance.py +++ b/airflow-core/tests/unit/models/test_taskinstance.py @@ -4108,8 +4108,28 @@ async def empty_callback_for_deadline(): pass -def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, session): - """Test that clearing tasks recalculates all (and only) DAGRUN_QUEUED_AT deadlines.""" +@pytest.mark.parametrize( + "use_variable_interval", + [ + pytest.param(False, id="fixed_timedelta_interval"), + pytest.param(True, id="variable_interval"), + ], +) +def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, session, use_variable_interval): + """Test that clearing tasks recalculates all (and only) DAGRUN_QUEUED_AT deadlines. + + Since Airflow 3.3 the ``deadline_alert.interval`` column is JSON (a serialized ``timedelta`` + or ``VariableInterval``), so the recalculation must decode it instead of passing the raw value + to ``timedelta()``. Storing the interval via ``serialize`` here mirrors production and covers + both interval kinds. + """ + from airflow.sdk.definitions.deadline import VariableInterval + from airflow.sdk.definitions.variable import Variable + from airflow.sdk.serde import serialize + + variable_key = "deadline_interval_key" + variable_seconds = 3600 + with dag_maker( dag_id="test_recalculate_deadlines", schedule=datetime.timedelta(days=1), @@ -4129,29 +4149,52 @@ def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, se select(SerializedDagModel.id).where(SerializedDagModel.dag_id == dag.dag_id) ) - deadline_configs = [ - (DeadlineReference.DAGRUN_QUEUED_AT, datetime.timedelta(hours=1)), - (DeadlineReference.DAGRUN_QUEUED_AT, datetime.timedelta(hours=2)), - (DeadlineReference.FIXED_DATETIME, datetime.timedelta(hours=1)), + if use_variable_interval: + queued_interval = VariableInterval(variable_key) + queued_resolved = datetime.timedelta(seconds=variable_seconds) + # Both DAGRUN_QUEUED_AT deadlines resolve to the same interval from the Variable. + queued_configs = [ + (DeadlineReference.DAGRUN_QUEUED_AT, queued_interval, queued_resolved), + (DeadlineReference.DAGRUN_QUEUED_AT, queued_interval, queued_resolved), + ] + else: + queued_configs = [ + ( + DeadlineReference.DAGRUN_QUEUED_AT, + datetime.timedelta(hours=1), + datetime.timedelta(hours=1), + ), + ( + DeadlineReference.DAGRUN_QUEUED_AT, + datetime.timedelta(hours=2), + datetime.timedelta(hours=2), + ), + ] + + # (reference type, interval stored on the alert, timedelta it resolves to) + deadline_configs = queued_configs + [ + (DeadlineReference.FIXED_DATETIME, datetime.timedelta(hours=1), datetime.timedelta(hours=1)), ] - for deadline_type, interval in deadline_configs: + expected_resolved_by_alert = {} + for deadline_type, stored_interval, resolved_interval in deadline_configs: if deadline_type == DeadlineReference.DAGRUN_QUEUED_AT: reference = DeadlineReference.DAGRUN_QUEUED_AT.serialize_reference() - deadline_time = dag_run.queued_at + interval + deadline_time = dag_run.queued_at + resolved_interval else: # FIXED_DATETIME future_date = timezone.utcnow() + datetime.timedelta(days=7) reference = DeadlineReference.FIXED_DATETIME(future_date).serialize_reference() - deadline_time = future_date + interval + deadline_time = future_date + resolved_interval deadline_alert = DeadlineAlertModel( serialized_dag_id=serialized_dag_id, reference=reference, - interval=interval.total_seconds(), + interval=serialize(stored_interval), callback_def={"path": f"{__name__}.empty_callback_for_deadline", "kwargs": {}}, ) session.add(deadline_alert) session.flush() + expected_resolved_by_alert[deadline_alert.id] = resolved_interval deadline = Deadline( dagrun_id=dag_run.id, @@ -4170,7 +4213,15 @@ def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, se } tis = session.scalars(select(TI).where(TI.dag_id == dag.dag_id, TI.run_id == dag_run.run_id)).all() - clear_task_instances(tis, session) + + # VariableInterval.resolve() reads the Airflow Variable during recalculation. + variable_ctx = ( + mock.patch.object(Variable, "get", return_value=str(variable_seconds)) + if use_variable_interval + else contextlib.nullcontext() + ) + with variable_ctx: + clear_task_instances(tis, session) dag_run = session.scalar(select(DagRun).where(DagRun.id == dag_run.id)) assert dag_run.queued_at > original_queued_at @@ -4183,8 +4234,7 @@ def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, se for deadline in deadlines_after: if deadline.deadline_time != deadline_times_by_alert[deadline.deadline_alert_id]: recalculated_count += 1 - deadline_alert = session.get(DeadlineAlertModel, deadline.deadline_alert_id) - expected_time = dag_run.queued_at + datetime.timedelta(seconds=deadline_alert.interval) + expected_time = dag_run.queued_at + expected_resolved_by_alert[deadline.deadline_alert_id] assert deadline.deadline_time == expected_time assert recalculated_count == 2 diff --git a/generated/known_sdk_imports_in_core.txt b/generated/known_sdk_imports_in_core.txt index f0aaae24ac227..6e98d3d3f4052 100644 --- a/generated/known_sdk_imports_in_core.txt +++ b/generated/known_sdk_imports_in_core.txt @@ -25,7 +25,7 @@ airflow-core/src/airflow/providers_manager.py::5 airflow-core/src/airflow/secrets/__init__.py::1 airflow-core/src/airflow/serialization/decoders.py::2 airflow-core/src/airflow/serialization/definitions/baseoperator.py::1 -airflow-core/src/airflow/serialization/definitions/dag.py::2 +airflow-core/src/airflow/serialization/definitions/dag.py::1 airflow-core/src/airflow/serialization/definitions/deadline.py::1 airflow-core/src/airflow/serialization/definitions/mappedoperator.py::5 airflow-core/src/airflow/serialization/encoders.py::11