From db15212be88a7585d9a89d78c8d0f189d830dea6 Mon Sep 17 00:00:00 2001 From: imrichardwu Date: Thu, 30 Jul 2026 18:38:48 -0700 Subject: [PATCH 1/8] Fix custom Deadline references with optional keyword arguments Preserve the Dag run evaluation context for custom references that accept additional keyword arguments, so they continue to evaluate correctly. --- airflow-core/docs/howto/deadline-alerts.rst | 20 +- .../newsfragments/70714.significant.rst | 1 + airflow-core/src/airflow/models/deadline.py | 126 +++------- .../src/airflow/models/taskinstance.py | 21 +- .../airflow/serialization/definitions/dag.py | 2 +- .../serialization/definitions/deadline.py | 87 ++----- .../tests/unit/models/test_deadline.py | 226 ++++-------------- .../tests/unit/models/test_deadline_alert.py | 79 +++++- .../tests/unit/models/test_taskinstance.py | 27 ++- .../src/airflow/sdk/definitions/deadline.py | 30 +-- 10 files changed, 221 insertions(+), 398 deletions(-) create mode 100644 airflow-core/newsfragments/70714.significant.rst diff --git a/airflow-core/docs/howto/deadline-alerts.rst b/airflow-core/docs/howto/deadline-alerts.rst index df09430862b3c..9fe4530b7e1df 100644 --- a/airflow-core/docs/howto/deadline-alerts.rst +++ b/airflow-core/docs/howto/deadline-alerts.rst @@ -437,9 +437,8 @@ implement an ``_evaluate_with()`` method. class MyCustomDecoratedReference(BaseDeadlineReference): """A custom reference evaluated when Dag runs are created.""" - def _evaluate_with(self, *, session: Session, **kwargs) -> datetime: - # Add your business logic here - return your_datetime + def _evaluate_with(self, *, session: Session, dagrun) -> datetime: + return dagrun.logical_date # You can specify when evaluate_with will be called by providing a DeadlineReference.TYPES value. @@ -447,14 +446,8 @@ implement an ``_evaluate_with()`` method. class MyQueuedReference(BaseDeadlineReference): """A custom reference evaluated when Dag runs are queued.""" - # Ask for the Dag run context values supplied by Airflow; see notes below. - required_kwargs = {"dag_id", "run_id"} - - def _evaluate_with(self, *, session: Session, **kwargs) -> datetime: - dag_id = kwargs["dag_id"] - run_id = kwargs["run_id"] - # Use dag_id and run_id in your calculation - return your_datetime + def _evaluate_with(self, *, session: Session, dagrun) -> datetime: + return dagrun.queued_at **Using a Custom Reference in a Dag** @@ -524,8 +517,5 @@ followed by a more urgent escalation if the Dag is still running. * **Timezone Awareness**: Always return timezone-aware datetime objects. * **Plugin Placement**: One convenient place for custom references is in the plugins directory. * **API Server Restart**: Restart the Airflow API Server after adding or modifying custom references. -* **Required Parameters**: ``required_kwargs`` declares which Dag run context values Airflow should - forward to ``_evaluate_with()``. Only ``dag_id`` and ``run_id`` are available; declaring anything - else raises a ``ValueError`` when the deadline is evaluated. To configure a reference itself, give - it constructor fields or read from an Airflow Variable. +* **Dag run context**: Add a ``dagrun`` parameter to ``_evaluate_with()`` to use the Dag run being evaluated. * **Database Access**: Use the ``session`` parameter for Airflow database queries if needed. diff --git a/airflow-core/newsfragments/70714.significant.rst b/airflow-core/newsfragments/70714.significant.rst new file mode 100644 index 0000000000000..70572b6d59845 --- /dev/null +++ b/airflow-core/newsfragments/70714.significant.rst @@ -0,0 +1 @@ +Deprecate the ``required_kwargs`` attribute for custom deadline references. Implement ``_evaluate_with()`` with a ``dagrun`` parameter instead. diff --git a/airflow-core/src/airflow/models/deadline.py b/airflow-core/src/airflow/models/deadline.py index 67438ce428abb..b35d09bad43a4 100644 --- a/airflow-core/src/airflow/models/deadline.py +++ b/airflow-core/src/airflow/models/deadline.py @@ -17,21 +17,23 @@ from __future__ import annotations import logging +import warnings from abc import ABC, abstractmethod from collections.abc import Sequence from dataclasses import dataclass from datetime import datetime, timedelta +from inspect import signature from typing import TYPE_CHECKING, Any, cast from uuid import UUID import uuid6 -from sqlalchemy import Boolean, ForeignKey, Index, Integer, Uuid, and_, func, inspect, select, text -from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy import Boolean, ForeignKey, Index, Integer, Uuid, and_, func, select, text from sqlalchemy.orm import Mapped, mapped_column, relationship from airflow._shared.observability.metrics import stats from airflow._shared.timezones import timezone from airflow.configuration import conf +from airflow.exceptions import RemovedInAirflow4Warning from airflow.models.base import Base from airflow.models.callback import ( Callback, @@ -57,6 +59,27 @@ CALLBACK_METRICS_PREFIX = "deadline_alerts" +def _get_evaluation_kwargs(reference: Any, evaluator: Any, kwargs: dict[str, Any]) -> dict[str, Any]: + """Return the evaluation arguments accepted by a deadline reference.""" + required_kwargs: set[str] | None = getattr(reference, "required_kwargs", None) + if required_kwargs is not None: + warnings.warn( + "required_kwargs is deprecated. Declare the keyword-only parameters your " + "_evaluate_with() implementation needs instead.", + RemovedInAirflow4Warning, + stacklevel=3, + ) + kwargs = {key: value for key, value in kwargs.items() if key in required_kwargs} + if missing_kwargs := required_kwargs - kwargs.keys(): + raise ValueError( + f"{reference.__class__.__name__} is missing required parameters: {', '.join(missing_kwargs)}" + ) + return kwargs + + parameters = signature(evaluator).parameters + return {key: value for key, value in kwargs.items() if key in parameters} + + class classproperty: """ Decorator that converts a method with a single cls argument into a property. @@ -320,34 +343,18 @@ def get_reference_class(cls, reference_name: str) -> type[BaseDeadlineReference] class BaseDeadlineReference(LoggingMixin, ABC): """Base class for all Deadline implementations.""" - # Set of required kwargs - subclasses should override this. - required_kwargs: set[str] = set() - @classproperty def reference_name(cls: Any) -> str: return cls.__name__ def evaluate_with(self, *, session: Session, interval: timedelta, **kwargs: Any) -> datetime | None: - """Validate the provided kwargs and evaluate this deadline with the given conditions.""" - filtered_kwargs = {k: v for k, v in kwargs.items() if k in self.required_kwargs} - - if missing_kwargs := self.required_kwargs - filtered_kwargs.keys(): - raise ValueError( - f"{self.__class__.__name__} is missing required parameters: {', '.join(missing_kwargs)}" - ) - - if extra_kwargs := kwargs.keys() - filtered_kwargs.keys(): - self.log.debug( - "%s ignoring unexpected parameters: %s", - self.reference_name, - ", ".join(extra_kwargs), - ) - - base_time = self._evaluate_with(session=session, **filtered_kwargs) + """Evaluate this deadline with the supplied context.""" + evaluation_kwargs = _get_evaluation_kwargs(self, self._evaluate_with, kwargs) + base_time = self._evaluate_with(session=session, **evaluation_kwargs) return base_time + interval if base_time is not None else None @abstractmethod - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: """Must be implemented by subclasses to perform the actual evaluation.""" raise NotImplementedError @@ -381,7 +388,7 @@ class FixedDatetimeDeadline(BaseDeadlineReference): _datetime: datetime - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: return self._datetime def serialize_reference(self) -> dict: @@ -397,23 +404,14 @@ def deserialize_reference(cls, reference_data: dict): class DagRunLogicalDateDeadline(BaseDeadlineReference): """A deadline that returns a DagRun's logical date.""" - required_kwargs = {"dag_id", "run_id"} - - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: - from airflow.models import DagRun - - return _fetch_from_db(DagRun.logical_date, session=session, **kwargs) + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + return dagrun.logical_date class DagRunQueuedAtDeadline(BaseDeadlineReference): """A deadline that returns when a DagRun was queued.""" - required_kwargs = {"dag_id", "run_id"} - - @provide_session - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: - from airflow.models import DagRun - - return _fetch_from_db(DagRun.queued_at, session=session, **kwargs) + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + return dagrun.queued_at @dataclass class AverageRuntimeDeadline(BaseDeadlineReference): @@ -422,7 +420,6 @@ class AverageRuntimeDeadline(BaseDeadlineReference): DEFAULT_LIMIT = 10 max_runs: int min_runs: int | None = None - required_kwargs = {"dag_id"} def __post_init__(self): if self.min_runs is None: @@ -431,10 +428,10 @@ def __post_init__(self): raise ValueError("min_runs must be at least 1") @provide_session - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: from airflow.models import DagRun - dag_id = kwargs["dag_id"] + dag_id = dagrun.dag_id # Get database dialect to use appropriate time difference calculation dialect = get_dialect_name(session) @@ -507,54 +504,3 @@ def deserialize_reference(cls, reference_data: dict): DeadlineReferenceType = ReferenceModels.BaseDeadlineReference - - -@provide_session -def _fetch_from_db(model_reference: Mapped, *, session=None, **conditions) -> datetime | None: - """ - Fetch a datetime value from the database using the provided model reference and filtering conditions. - - For example, to fetch a TaskInstance's start_date: - _fetch_from_db( - TaskInstance.start_date, dag_id='example_dag', task_id='example_task', run_id='example_run' - ) - - This generates SQL equivalent to: - SELECT start_date - FROM task_instance - WHERE dag_id = 'example_dag' - AND task_id = 'example_task' - AND run_id = 'example_run' - - :param model_reference: SQLAlchemy Column to select (e.g., DagRun.logical_date, TaskInstance.start_date) - :param conditions: Filtering conditions applied as equality comparisons in the WHERE clause. - Multiple conditions are combined with AND. - :param session: SQLAlchemy session (auto-provided by decorator) - """ - query = select(model_reference) - - for key, value in conditions.items(): - inspected = inspect(model_reference) - if inspected is not None: - query = query.where(getattr(inspected.class_, key) == value) - - compiled_query = query.compile(compile_kwargs={"literal_binds": True}) - pretty_query = "\n ".join(str(compiled_query).splitlines()) - logger.debug( - "Executing query:\n %r\nAs SQL:\n %s", - query, - pretty_query, - ) - - try: - result = session.scalar(query) - except SQLAlchemyError: - logger.exception("Database query failed.") - raise - - if result is None: - message = f"No matching record found in the database for query:\n {pretty_query}" - logger.error(message) - raise ValueError(message) - - return result diff --git a/airflow-core/src/airflow/models/taskinstance.py b/airflow-core/src/airflow/models/taskinstance.py index 0dd10507fc863..59b9ad3baea24 100644 --- a/airflow-core/src/airflow/models/taskinstance.py +++ b/airflow-core/src/airflow/models/taskinstance.py @@ -94,6 +94,7 @@ from airflow.models.taskmap import TaskMap from airflow.models.taskreschedule import TaskReschedule from airflow.models.xcom import XCOM_RETURN_KEY, LazyXComSelectSequence, XComModel +from airflow.serialization.decoders import decode_deadline_reference from airflow.serialization.enums import stringify_encoding_keys from airflow.settings import task_instance_mutation_hook from airflow.task.priority_strategy import validate_and_load_priority_weight_strategy @@ -222,14 +223,11 @@ def _add_and_prime_mapped_ti( set_committed_value(ti, "dag_run", dag_run) -def _recalculate_dagrun_queued_at_deadlines( - dagrun: DagRun, new_queued_at: datetime, session: Session -) -> None: +def _recalculate_dagrun_queued_at_deadlines(dagrun: DagRun, *, session: Session) -> None: """ Recalculate deadline times for deadlines that reference dagrun.queued_at. :param dagrun: The DagRun whose deadlines should be recalculated - :param new_queued_at: The new queued_at timestamp to use for calculation :param session: Database session :meta private: @@ -249,9 +247,16 @@ def _recalculate_dagrun_queued_at_deadlines( return 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 + new_deadline_time = decode_deadline_reference(deadline_alert.reference).evaluate_with( + session=session, + interval=timedelta(seconds=deadline_alert.interval), + dagrun=dagrun, + dag_id=dagrun.dag_id, + run_id=dagrun.run_id, + ) + + if new_deadline_time is None: + continue log.debug( "Recalculating deadline %s for DagRun %s.%s: old=%s, new=%s", @@ -459,7 +464,7 @@ def clear_task_instances( parent_context=parent_trace_context(dr.conf), ) - _recalculate_dagrun_queued_at_deadlines(dr, dr.queued_at, session) + _recalculate_dagrun_queued_at_deadlines(dr, session=session) if dr.state in State.finished_dr_states: dr.state = dag_run_state diff --git a/airflow-core/src/airflow/serialization/definitions/dag.py b/airflow-core/src/airflow/serialization/definitions/dag.py index 8ed0fee2ccabd..b3d9be2807f70 100644 --- a/airflow-core/src/airflow/serialization/definitions/dag.py +++ b/airflow-core/src/airflow/serialization/definitions/dag.py @@ -761,7 +761,7 @@ def _process_dagrun_deadline_alerts( deadline_time = deserialized_deadline_alert.reference.evaluate_with( session=session, interval=interval, - # TODO : Pretty sure we can drop these last two; verify after testing is complete + dagrun=orm_dagrun, dag_id=self.dag_id, run_id=orm_dagrun.run_id, ) diff --git a/airflow-core/src/airflow/serialization/definitions/deadline.py b/airflow-core/src/airflow/serialization/definitions/deadline.py index 20fac2b54e874..f39112d831733 100644 --- a/airflow-core/src/airflow/serialization/definitions/deadline.py +++ b/airflow-core/src/airflow/serialization/definitions/deadline.py @@ -27,7 +27,7 @@ from sqlalchemy import select from airflow._shared.timezones import timezone -from airflow.models.deadline import classproperty +from airflow.models.deadline import _get_evaluation_kwargs, classproperty from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.session import provide_session from airflow.utils.sqlalchemy import get_dialect_name @@ -96,33 +96,18 @@ def get_reference_class(cls, reference_name: str) -> type[SerializedBaseDeadline class SerializedBaseDeadlineReference(LoggingMixin, ABC): """Base class for all serialized Deadline implementations.""" - required_kwargs: set[str] = set() - @classproperty def reference_name(cls: Any) -> str: return cls.__name__ def evaluate_with(self, *, session: Session, interval: timedelta, **kwargs: Any) -> datetime | None: - """Validate the provided kwargs and evaluate this deadline with the given conditions.""" - filtered_kwargs = {k: v for k, v in kwargs.items() if k in self.required_kwargs} - - if missing_kwargs := self.required_kwargs - filtered_kwargs.keys(): - raise ValueError( - f"{self.__class__.__name__} is missing required parameters: {', '.join(missing_kwargs)}" - ) - - if extra_kwargs := kwargs.keys() - filtered_kwargs.keys(): - self.log.debug( - "%s ignoring unexpected parameters: %s", - self.reference_name, - ", ".join(extra_kwargs), - ) - - base_time = self._evaluate_with(session=session, **filtered_kwargs) + """Evaluate this deadline with the supplied context.""" + evaluation_kwargs = _get_evaluation_kwargs(self, self._evaluate_with, kwargs) + base_time = self._evaluate_with(session=session, **evaluation_kwargs) return base_time + interval if base_time is not None else None @abstractmethod - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: """Must be implemented by subclasses to perform the actual evaluation.""" raise NotImplementedError @@ -153,7 +138,7 @@ class FixedDatetimeDeadline(SerializedBaseDeadlineReference): _datetime: datetime - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: return self._datetime def serialize_reference(self) -> dict: @@ -169,23 +154,14 @@ def deserialize_reference(cls, reference_data: dict): class DagRunLogicalDateDeadline(SerializedBaseDeadlineReference): """A deadline that returns a DagRun's logical date.""" - required_kwargs = {"dag_id", "run_id"} - - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: - from airflow.models import DagRun - - return _fetch_from_db(DagRun.logical_date, session=session, **kwargs) + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + return dagrun.logical_date class DagRunQueuedAtDeadline(SerializedBaseDeadlineReference): """A deadline that returns when a DagRun was queued.""" - required_kwargs = {"dag_id", "run_id"} - - @provide_session - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: - from airflow.models import DagRun - - return _fetch_from_db(DagRun.queued_at, session=session, **kwargs) + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + return dagrun.queued_at @dataclass class AverageRuntimeDeadline(SerializedBaseDeadlineReference): @@ -194,7 +170,6 @@ class AverageRuntimeDeadline(SerializedBaseDeadlineReference): DEFAULT_LIMIT = 10 max_runs: int min_runs: int | None = None - required_kwargs = {"dag_id"} def __post_init__(self): if self.min_runs is None: @@ -203,13 +178,13 @@ def __post_init__(self): raise ValueError("min_runs must be at least 1") @provide_session - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: - from sqlalchemy import func, select, text + def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + from sqlalchemy import func, text from airflow.models import DagRun from airflow.utils.state import DagRunState - dag_id = kwargs["dag_id"] + dag_id = dagrun.dag_id dialect = get_dialect_name(session) @@ -280,7 +255,7 @@ class SerializedCustomReference(SerializedBaseDeadlineReference): """ Wrapper for custom deadline references. - This class dynamically delegates to the wrapped reference for required_kwargs and evaluation logic. + This class dynamically delegates to the wrapped reference's evaluation logic. """ def __init__(self, inner_ref): @@ -291,23 +266,9 @@ def reference_name(self) -> str: return self.inner_ref.reference_name def evaluate_with(self, *, session: Session, interval: timedelta, **kwargs: Any) -> datetime | None: - """Validate the provided kwargs and evaluate this deadline with the given conditions.""" - required_kwargs: set[str] = getattr(self.inner_ref, "required_kwargs", set()) - filtered_kwargs = {k: v for k, v in kwargs.items() if k in required_kwargs} - - if missing_kwargs := required_kwargs - filtered_kwargs.keys(): - raise ValueError( - f"{self.inner_ref.__class__.__name__} is missing required parameters: {', '.join(missing_kwargs)}" - ) - - if extra_kwargs := kwargs.keys() - filtered_kwargs.keys(): - self.log.debug( - "%s ignoring unexpected parameters: %s", - self.reference_name, - ", ".join(extra_kwargs), - ) - - deadline = self.inner_ref._evaluate_with(session=session, **filtered_kwargs) + """Evaluate the wrapped custom reference with the supplied context.""" + evaluation_kwargs = _get_evaluation_kwargs(self.inner_ref, self.inner_ref._evaluate_with, kwargs) + deadline = self.inner_ref._evaluate_with(session=session, **evaluation_kwargs) return deadline + interval if deadline is not None else None def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: @@ -358,20 +319,6 @@ class TYPES: ) -def _fetch_from_db(column, *, session: Session, dag_id: str, run_id: str) -> datetime | None: - """ - Fetch a datetime column from the DagRun table. - - :meta private: - """ - from airflow.models import DagRun - - result = session.execute(select(column).where(DagRun.dag_id == dag_id, DagRun.run_id == run_id)).scalar() - if result is None: - logger.warning("Could not find DagRun for dag_id=%s, run_id=%s", dag_id, run_id) - return result - - @attrs.define class SerializedDeadlineAlert: """Serialized representation of a deadline alert.""" diff --git a/airflow-core/tests/unit/models/test_deadline.py b/airflow-core/tests/unit/models/test_deadline.py index 51e9d208c07f4..7a74fe8e92d74 100644 --- a/airflow-core/tests/unit/models/test_deadline.py +++ b/airflow-core/tests/unit/models/test_deadline.py @@ -18,17 +18,17 @@ import re from datetime import datetime, timedelta +from types import SimpleNamespace from typing import TYPE_CHECKING from unittest import mock import pytest import time_machine from sqlalchemy import select -from sqlalchemy.exc import SQLAlchemyError from airflow.api_fastapi.core_api.datamodels.dag_run import DAGRunResponse from airflow.models import DagRun -from airflow.models.deadline import Deadline, _fetch_from_db +from airflow.models.deadline import Deadline, ReferenceModels from airflow.providers.standard.operators.empty import EmptyOperator from airflow.sdk import timezone from airflow.sdk.definitions.callback import AsyncCallback, SyncCallback @@ -54,15 +54,6 @@ INVALID_DAG_ID = "invalid_dag_id" INVALID_RUN_ID = -1 -REFERENCE_TYPES = [ - pytest.param(SerializedReferenceModels.DagRunLogicalDateDeadline(), id="logical_date"), - pytest.param(SerializedReferenceModels.DagRunQueuedAtDeadline(), id="queued_at"), - pytest.param(SerializedReferenceModels.FixedDatetimeDeadline(DEFAULT_DATE), id="fixed_deadline"), - pytest.param( - SerializedReferenceModels.AverageRuntimeDeadline(max_runs=10, min_runs=10), id="average_runtime" - ), -] - async def callback_for_deadline(): """Used in a number of tests to confirm that Deadlines and DeadlineAlerts function correctly.""" @@ -317,138 +308,41 @@ def test_handle_miss_persists_executor_callback_routing_data(self, dagrun, sessi @pytest.mark.db_test -class TestCalculatedDeadlineDatabaseCalls: +class TestCalculatedDeadlineReferences: @staticmethod def teardown_method(): _clean_db() @pytest.mark.parametrize( - ("column", "conditions", "expected_query"), + ("reference", "attribute"), [ pytest.param( - DagRun.logical_date, - {"dag_id": DAG_ID}, - "SELECT dag_run.logical_date \nFROM dag_run \nWHERE dag_run.dag_id = :dag_id_1", - id="single_condition_logical_date", - ), - pytest.param( - DagRun.queued_at, - {"dag_id": DAG_ID}, - "SELECT dag_run.queued_at \nFROM dag_run \nWHERE dag_run.dag_id = :dag_id_1", - id="single_condition_queued_at", + SerializedReferenceModels.DagRunLogicalDateDeadline(), "logical_date", id="logical_date" ), + pytest.param(SerializedReferenceModels.DagRunQueuedAtDeadline(), "queued_at", id="queued_at"), pytest.param( - DagRun.logical_date, - {"dag_id": DAG_ID, "state": "running"}, - "SELECT dag_run.logical_date \nFROM dag_run \nWHERE dag_run.dag_id = :dag_id_1 AND dag_run.state = :state_1", - id="multiple_conditions", + ReferenceModels.DagRunLogicalDateDeadline(), "logical_date", id="legacy_logical_date" ), + pytest.param(ReferenceModels.DagRunQueuedAtDeadline(), "queued_at", id="legacy_queued_at"), ], ) - @mock.patch("sqlalchemy.orm.Session") - def test_fetch_from_db_success(self, mock_session, column, conditions, expected_query): - """Test successful database queries.""" - mock_session.scalar.return_value = DEFAULT_DATE - - result = _fetch_from_db(column, session=mock_session, **conditions) - - assert isinstance(result, datetime) - mock_session.scalar.assert_called_once() - - # Check that the correct query was constructed - call_args = mock_session.scalar.call_args[0][0] - assert str(call_args) == expected_query - - # Verify the actual parameter values - compiled = call_args.compile() - for key, value in conditions.items(): - # Note that SQLAlchemy appends the _1 to ensure unique template field names - assert compiled.params[f"{key}_1"] == value - - @pytest.mark.parametrize( - ("use_valid_conditions", "scalar_side_effect", "expected_error", "expected_message"), - [ - pytest.param( - False, - mock.DEFAULT, # This will allow the call to pass through - AttributeError, - None, - id="invalid_attribute", - ), - pytest.param( - True, - SQLAlchemyError("Database connection failed"), - SQLAlchemyError, - "Database connection failed", - id="database_error", - ), - pytest.param( - True, lambda x: None, ValueError, "No matching record found in the database", id="no_results" - ), - ], - ) - @mock.patch("sqlalchemy.orm.Session") - def test_fetch_from_db_error_cases( - self, mock_session, use_valid_conditions, scalar_side_effect, expected_error, expected_message - ): - """Test database access error handling.""" - model_reference = DagRun.logical_date - conditions = {"dag_id": "test_dag"} if use_valid_conditions else {"non_existent_column": "some_value"} - - # Configure mock session - mock_session.scalar.side_effect = scalar_side_effect - - with pytest.raises(expected_error, match=expected_message): - _fetch_from_db(model_reference, session=mock_session, **conditions) - - @pytest.mark.parametrize( - ("reference", "expected_column"), - [ - pytest.param( - SerializedReferenceModels.DagRunLogicalDateDeadline(), DagRun.logical_date, id="logical_date" - ), - pytest.param( - SerializedReferenceModels.DagRunQueuedAtDeadline(), DagRun.queued_at, id="queued_at" - ), - pytest.param( - SerializedReferenceModels.FixedDatetimeDeadline(DEFAULT_DATE), None, id="fixed_deadline" - ), - pytest.param( - SerializedReferenceModels.AverageRuntimeDeadline(max_runs=10, min_runs=10), - None, - id="average_runtime", - ), - ], - ) - def test_deadline_database_integration(self, reference, expected_column, session): - """ - Test database integration for all deadline types. - - Verifies: - 1. Calculated deadlines call _fetch_from_db with correct column. - 2. Fixed deadlines do not interact with database. - 3. Intervals are added to reference times. - """ - conditions = {"dag_id": DAG_ID, "run_id": "dagrun_1"} + def test_dagrun_references_use_supplied_dagrun(self, reference, attribute, session): + """DagRun references use the in-memory DagRun instead of querying it again.""" + dagrun = SimpleNamespace(logical_date=DEFAULT_DATE, queued_at=DEFAULT_DATE) interval = timedelta(hours=1) - with mock.patch("airflow.serialization.definitions.deadline._fetch_from_db") as mock_fetch: - mock_fetch.return_value = DEFAULT_DATE - - if expected_column is not None: - result = reference.evaluate_with(session=session, interval=interval, **conditions) - mock_fetch.assert_called_once_with(expected_column, session=session, **conditions) - elif isinstance(reference, SerializedReferenceModels.AverageRuntimeDeadline): - with mock.patch("airflow._shared.timezones.timezone.utcnow") as mock_utcnow: - mock_utcnow.return_value = DEFAULT_DATE - # No DAG runs exist, so it should use 24-hour default - result = reference.evaluate_with(session=session, interval=interval, dag_id=DAG_ID) - mock_fetch.assert_not_called() - # Should return None when no DAG runs exist - assert result is None - else: - result = reference.evaluate_with(session=session, interval=interval) - mock_fetch.assert_not_called() - assert result == DEFAULT_DATE + interval + + assert getattr(dagrun, attribute) == DEFAULT_DATE + assert ( + reference.evaluate_with( + session=session, + interval=interval, + dagrun=dagrun, + dag_id=DAG_ID, + run_id="dagrun_1", + unexpected="ignored", + ) + == DEFAULT_DATE + interval + ) def test_average_runtime_with_sufficient_history(self, session, dag_maker): """Test AverageRuntimeDeadline when enough historical data exists.""" @@ -479,7 +373,9 @@ def test_average_runtime_with_sufficient_history(self, session, dag_maker): with mock.patch("airflow._shared.timezones.timezone.utcnow") as mock_utcnow: mock_utcnow.return_value = DEFAULT_DATE - result = reference.evaluate_with(session=session, interval=interval, dag_id=DAG_ID) + result = reference.evaluate_with( + session=session, interval=interval, dagrun=SimpleNamespace(dag_id=DAG_ID) + ) # Calculate expected average: sum(durations) / len(durations) expected_avg_seconds = sum(durations) / len(durations) @@ -516,7 +412,9 @@ def test_average_runtime_with_insufficient_history(self, session, dag_maker): with mock.patch("airflow._shared.timezones.timezone.utcnow") as mock_utcnow: mock_utcnow.return_value = DEFAULT_DATE - result = reference.evaluate_with(session=session, interval=interval, dag_id=DAG_ID) + result = reference.evaluate_with( + session=session, interval=interval, dagrun=SimpleNamespace(dag_id=DAG_ID) + ) # Should return None since insufficient runs assert result is None @@ -550,7 +448,9 @@ def test_average_runtime_with_min_runs(self, session, dag_maker): with mock.patch("airflow._shared.timezones.timezone.utcnow") as mock_utcnow: mock_utcnow.return_value = DEFAULT_DATE - result = reference.evaluate_with(session=session, interval=interval, dag_id=DAG_ID) + result = reference.evaluate_with( + session=session, interval=interval, dagrun=SimpleNamespace(dag_id=DAG_ID) + ) # Should calculate average from 3 runs expected_avg_seconds = sum(durations) / len(durations) # 4200 seconds @@ -563,7 +463,9 @@ def test_average_runtime_with_min_runs(self, session, dag_maker): with mock.patch("airflow._shared.timezones.timezone.utcnow") as mock_utcnow: mock_utcnow.return_value = DEFAULT_DATE - result = reference.evaluate_with(session=session, interval=interval, dag_id=DAG_ID) + result = reference.evaluate_with( + session=session, interval=interval, dagrun=SimpleNamespace(dag_id=DAG_ID) + ) assert result is None def test_average_runtime_min_runs_validation(self): @@ -612,7 +514,9 @@ def test_average_runtime_excludes_non_successful_runs(self, session, dag_maker): with mock.patch("airflow._shared.timezones.timezone.utcnow") as mock_utcnow: mock_utcnow.return_value = DEFAULT_DATE - result = reference.evaluate_with(session=session, interval=interval, dag_id=DAG_ID) + result = reference.evaluate_with( + session=session, interval=interval, dagrun=SimpleNamespace(dag_id=DAG_ID) + ) # Average must be over the 3 successful 60s runs only (not the 36000s failures). expected = DEFAULT_DATE + timedelta(seconds=success_duration) + interval @@ -636,61 +540,15 @@ def test_average_runtime_skips_when_too_few_successful_runs(self, session, dag_m session.commit() reference = SerializedReferenceModels.AverageRuntimeDeadline(max_runs=10, min_runs=3) - result = reference.evaluate_with(session=session, interval=timedelta(hours=1), dag_id=DAG_ID) + result = reference.evaluate_with( + session=session, interval=timedelta(hours=1), dagrun=SimpleNamespace(dag_id=DAG_ID) + ) assert result is None class TestDeadlineReference: """DeadlineReference lives in definitions/deadlines.py but properly testing them requires DB access.""" - DEFAULT_INTERVAL = timedelta(hours=1) - DEFAULT_ARGS = {"interval": DEFAULT_INTERVAL} - - @pytest.mark.parametrize("reference", REFERENCE_TYPES) - @pytest.mark.db_test - def test_deadline_evaluate_with(self, reference, session): - """Test that all deadline types evaluate correctly with their required conditions.""" - conditions = { - "dag_id": DAG_ID, - "run_id": "dagrun_1", - "unexpected": "param", # Add an unexpected parameter. - "extra": "kwarg", # Add another unexpected parameter. - } - - with mock.patch.object(reference, "_evaluate_with") as mock_evaluate: - mock_evaluate.return_value = DEFAULT_DATE - - if reference.required_kwargs: - result = reference.evaluate_with(**self.DEFAULT_ARGS, session=session, **conditions) - else: - result = reference.evaluate_with(**self.DEFAULT_ARGS, session=session) - - # Verify only expected kwargs are passed through. - expected_kwargs = {k: conditions[k] for k in reference.required_kwargs if k in conditions} - expected_kwargs["session"] = session - - mock_evaluate.assert_called_once_with(**expected_kwargs) - - assert result == DEFAULT_DATE + self.DEFAULT_INTERVAL - - @pytest.mark.parametrize("reference", REFERENCE_TYPES) - @pytest.mark.db_test - def test_deadline_missing_required_kwargs(self, reference, session): - """Test that deadlines raise appropriate errors for missing required parameters.""" - if reference.required_kwargs: - with pytest.raises( - ValueError, match=re.escape(f"{reference.__class__.__name__} is missing required parameters:") - ) as raised_exception: - reference.evaluate_with(session=session, **self.DEFAULT_ARGS) - - assert all(substring in str(raised_exception.value) for substring in reference.required_kwargs) - else: - # Let the lack of an exception here effectively assert that no exception is raised. - reference.evaluate_with(session=session, **self.DEFAULT_ARGS) - - for required_param in reference.required_kwargs: - assert required_param in str(raised_exception.value) - def test_deadline_reference_creation(self): """Test that DeadlineReference provides consistent interface and types.""" fixed_reference = DeadlineReference.FIXED_DATETIME(DEFAULT_DATE) diff --git a/airflow-core/tests/unit/models/test_deadline_alert.py b/airflow-core/tests/unit/models/test_deadline_alert.py index a9b1854f6abce..935f71beeb60e 100644 --- a/airflow-core/tests/unit/models/test_deadline_alert.py +++ b/airflow-core/tests/unit/models/test_deadline_alert.py @@ -17,6 +17,7 @@ from __future__ import annotations from datetime import timedelta +from types import SimpleNamespace from unittest.mock import Mock import pytest @@ -24,6 +25,7 @@ from sqlalchemy import select from airflow._shared.timezones import timezone +from airflow.exceptions import RemovedInAirflow4Warning from airflow.models.deadline_alert import DeadlineAlert from airflow.models.serialized_dag import SerializedDagModel from airflow.sdk.definitions.deadline import BaseDeadlineReference, DeadlineReference @@ -188,23 +190,78 @@ def _evaluate_with(self, *, session, dag_id, run_id): wrapper = SerializedReferenceModels.SerializedCustomReference(inner_ref) - wrapper.evaluate_with( + with pytest.warns(RemovedInAirflow4Warning, match="required_kwargs is deprecated"): + wrapper.evaluate_with( + session=None, + interval=timedelta(hours=1), + dag_id="test_dag", + run_id="test_run", + extra_param="should_be_filtered", + ) + + inner_ref._evaluate_with.assert_called_once_with(session=None, dag_id="test_dag", run_id="test_run") + + # try calling with missing required parameters + with pytest.warns(RemovedInAirflow4Warning, match="required_kwargs is deprecated"): + with pytest.raises(ValueError, match="missing required parameters: run_id"): + wrapper.evaluate_with( + session=None, + interval=timedelta(hours=1), + dag_id="test_dag", + ) + + def test_serialized_custom_reference_receives_dagrun(self): + class DagRunReference(BaseDeadlineReference): + def _evaluate_with(self, *, session, dagrun): + return dagrun.queued_at + + wrapper = SerializedReferenceModels.SerializedCustomReference(DagRunReference()) + dagrun = SimpleNamespace(queued_at=DEFAULT_DATE) + + assert wrapper.evaluate_with( session=None, interval=timedelta(hours=1), + dagrun=dagrun, dag_id="test_dag", run_id="test_run", - extra_param="should_be_filtered", - ) + unexpected="ignored", + ) == DEFAULT_DATE + timedelta(hours=1) - inner_ref._evaluate_with.assert_called_once_with(session=None, dag_id="test_dag", run_id="test_run") + def test_serialized_custom_reference_receives_dagrun_with_var_keyword_arguments(self): + class DagRunReference(BaseDeadlineReference): + def _evaluate_with(self, *, session, dagrun, **kwargs): + return dagrun.queued_at - # try calling with missing required parameters - with pytest.raises(ValueError, match="missing required parameters: run_id"): - wrapper.evaluate_with( - session=None, - interval=timedelta(hours=1), - dag_id="test_dag", - ) + wrapper = SerializedReferenceModels.SerializedCustomReference(DagRunReference()) + dagrun = SimpleNamespace(queued_at=DEFAULT_DATE) + + assert wrapper.evaluate_with( + session=None, + interval=timedelta(hours=1), + dagrun=dagrun, + dag_id="test_dag", + run_id="test_run", + ) == DEFAULT_DATE + timedelta(hours=1) + + def test_legacy_serialized_custom_reference_ignores_evaluation_context(self): + class LegacyReference(BaseDeadlineReference): + received_kwargs: dict[str, object] + + def _evaluate_with(self, *, session, **kwargs): + self.received_kwargs = kwargs + return DEFAULT_DATE + + reference = LegacyReference() + wrapper = SerializedReferenceModels.SerializedCustomReference(reference) + + assert wrapper.evaluate_with( + session=None, + interval=timedelta(hours=1), + dagrun=SimpleNamespace(), + dag_id="test_dag", + run_id="test_run", + ) == DEFAULT_DATE + timedelta(hours=1) + assert reference.received_kwargs == {} def test_core_deadline_reference_treated_as_builtins(self): """Test that refs from airflow.models.deadline are still treated as builtins.""" diff --git a/airflow-core/tests/unit/models/test_taskinstance.py b/airflow-core/tests/unit/models/test_taskinstance.py index 53fa99aa22460..a6af419f55ab2 100644 --- a/airflow-core/tests/unit/models/test_taskinstance.py +++ b/airflow-core/tests/unit/models/test_taskinstance.py @@ -4107,8 +4107,28 @@ async def empty_callback_for_deadline(): pass -def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, session): +def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, session, monkeypatch): """Test that clearing tasks recalculates all (and only) DAGRUN_QUEUED_AT deadlines.""" + evaluation_calls = [] + + class QueuedDeadlineReference: + def evaluate_with(self, *, session, interval, dagrun, dag_id, run_id): + evaluation_calls.append( + { + "session": session, + "interval": interval, + "dagrun": dagrun, + "dag_id": dag_id, + "run_id": run_id, + } + ) + return dagrun.queued_at + interval + + monkeypatch.setattr( + "airflow.models.taskinstance.decode_deadline_reference", + lambda reference: QueuedDeadlineReference(), + ) + with dag_maker( dag_id="test_recalculate_deadlines", schedule=datetime.timedelta(days=1), @@ -4187,6 +4207,11 @@ def test_clear_task_instances_recalculates_dagrun_queued_deadlines(dag_maker, se assert deadline.deadline_time == expected_time assert recalculated_count == 2 + assert len(evaluation_calls) == 2 + assert all(call["session"] is session for call in evaluation_calls) + assert all(call["dagrun"] is dag_run for call in evaluation_calls) + assert all(call["dag_id"] == dag_run.dag_id for call in evaluation_calls) + assert all(call["run_id"] == dag_run.run_id for call in evaluation_calls) def test_get_dagrun_loaded_but_none_returns_dagrun(dag_maker, session): diff --git a/task-sdk/src/airflow/sdk/definitions/deadline.py b/task-sdk/src/airflow/sdk/definitions/deadline.py index a9da3dd3d3ea2..b751ee3367d0a 100644 --- a/task-sdk/src/airflow/sdk/definitions/deadline.py +++ b/task-sdk/src/airflow/sdk/definitions/deadline.py @@ -46,7 +46,8 @@ class BaseDeadlineReference(ABC): The actual evaluation logic (_evaluate_with) is in Core's SerializedReferenceModels. For custom deadline references, users should inherit from this class and implement - _evaluate_with() with deferred Core imports (imports inside the method body). + _evaluate_with() with deferred Core imports (imports inside the method body). Its + keyword-only ``dagrun`` parameter receives the DagRun being evaluated. """ @property @@ -183,8 +184,8 @@ class DeadlineReference: The public interface class for all DeadlineReference options. This class provides a unified interface for working with Deadlines, supporting both - calculated deadlines (which fetch values from the database) and fixed deadlines - (which return a predefined datetime). + calculated deadlines, including references to the DagRun being evaluated and + historical runtime data, and fixed deadlines (which return a predefined datetime). ------ Usage: @@ -211,17 +212,13 @@ class DeadlineReference: ), ) - 3. Evaluating deadlines will ignore unexpected parameters: + 3. Custom references receive the DagRun being evaluated: .. code-block:: python - # For deadlines requiring parameters: - deadline = DeadlineReference.DAGRUN_LOGICAL_DATE - deadline.evaluate_with(dag_id=dag.dag_id) - - # For deadlines with no required parameters: - deadline = DeadlineReference.FIXED_DATETIME(datetime(2025, 5, 4)) - deadline.evaluate_with() + class MyDeadlineReference(BaseDeadlineReference): + def _evaluate_with(self, *, session, dagrun): + return dagrun.logical_date """ class TYPES: @@ -320,10 +317,8 @@ def deadline_reference( @deadline_reference() class MyCustomReference(BaseDeadlineReference): # By default, evaluate_with will be called when a new dagrun is created. - def _evaluate_with(self, *, session: Session, **kwargs) -> datetime: - # Put your business logic here (use deferred imports for Core types) - from airflow.models import DagRun - return some_datetime + def _evaluate_with(self, *, session: Session, dagrun) -> datetime: + return dagrun.logical_date def serialize_reference(self) -> dict: return {"reference_type": self.reference_name} @@ -331,9 +326,8 @@ def serialize_reference(self) -> dict: @deadline_reference(DeadlineReference.TYPES.DAGRUN_QUEUED) class MyQueuedRef(BaseDeadlineReference): # Optionally, you can specify when you want it calculated by providing a DeadlineReference.TYPES - def _evaluate_with(self, *, session: Session, **kwargs) -> datetime: - # Put your business logic here - return some_datetime + def _evaluate_with(self, *, session: Session, dagrun) -> datetime: + return dagrun.queued_at def serialize_reference(self) -> dict: return {"reference_type": self.reference_name} From 4472afc1e80232f002fa56d87d43dc9ab25502be Mon Sep 17 00:00:00 2001 From: Richard Date: Thu, 6 Aug 2026 20:55:22 -0700 Subject: [PATCH 2/8] Update airflow-core/docs/howto/deadline-alerts.rst Co-authored-by: D. Ferruzzi --- airflow-core/docs/howto/deadline-alerts.rst | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/airflow-core/docs/howto/deadline-alerts.rst b/airflow-core/docs/howto/deadline-alerts.rst index 9d170a9f70cd5..5ea26f5c42780 100644 --- a/airflow-core/docs/howto/deadline-alerts.rst +++ b/airflow-core/docs/howto/deadline-alerts.rst @@ -442,7 +442,8 @@ choose a different time. """A custom reference evaluated when Dag runs are created.""" def _evaluate_with(self, *, session: Session, dagrun) -> datetime: - return dagrun.logical_date + my_datetime = my_business_logic(dagrun.logical_date) + return my_datetime # You can specify when evaluate_with will be called by providing a DeadlineReference.TYPES value. From 718b6607bea8eae55f38826cd7c49807a8bae33a Mon Sep 17 00:00:00 2001 From: Richard Date: Thu, 6 Aug 2026 20:55:29 -0700 Subject: [PATCH 3/8] Update airflow-core/src/airflow/models/deadline.py Co-authored-by: D. Ferruzzi --- airflow-core/src/airflow/models/deadline.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow-core/src/airflow/models/deadline.py b/airflow-core/src/airflow/models/deadline.py index b35d09bad43a4..c461b7b6bd86d 100644 --- a/airflow-core/src/airflow/models/deadline.py +++ b/airflow-core/src/airflow/models/deadline.py @@ -388,7 +388,7 @@ class FixedDatetimeDeadline(BaseDeadlineReference): _datetime: datetime - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, **kwargs) -> datetime | None: return self._datetime def serialize_reference(self) -> dict: From 4df9e257700d156af8ba477f21b0ea8f7a4c97ee Mon Sep 17 00:00:00 2001 From: Richard Date: Thu, 6 Aug 2026 20:56:17 -0700 Subject: [PATCH 4/8] Update airflow-core/src/airflow/models/deadline.py Co-authored-by: D. Ferruzzi --- airflow-core/src/airflow/models/deadline.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/airflow-core/src/airflow/models/deadline.py b/airflow-core/src/airflow/models/deadline.py index c461b7b6bd86d..332a8b6541142 100644 --- a/airflow-core/src/airflow/models/deadline.py +++ b/airflow-core/src/airflow/models/deadline.py @@ -77,6 +77,9 @@ def _get_evaluation_kwargs(reference: Any, evaluator: Any, kwargs: dict[str, Any return kwargs parameters = signature(evaluator).parameters + if any(param.kind is param.VAR_KEYWORD for param in parameters.values()): + return kwargs + return {key: value for key, value in kwargs.items() if key in parameters} From 608c2436c6ff9b0043d40cf52e96adf2ba88c54c Mon Sep 17 00:00:00 2001 From: imrichardwu Date: Thu, 6 Aug 2026 21:13:57 -0700 Subject: [PATCH 5/8] Type deadline reference dagrun as DagRunProtocol The custom deadline reference examples returned the dagrun attribute directly, which read as if the attribute was the required return value rather than an input available to the author's business logic. Type the evaluation dagrun with the shared DagRunProtocol so authors get a concrete contract for what is available, and show the attribute feeding into business logic in the examples. --- airflow-core/docs/howto/deadline-alerts.rst | 3 ++- airflow-core/src/airflow/models/deadline.py | 9 +++++---- task-sdk/src/airflow/sdk/definitions/deadline.py | 12 +++++++++--- uv.lock | 6 +++--- 4 files changed, 19 insertions(+), 11 deletions(-) diff --git a/airflow-core/docs/howto/deadline-alerts.rst b/airflow-core/docs/howto/deadline-alerts.rst index 0acadbc492f8e..5dc07759f92af 100644 --- a/airflow-core/docs/howto/deadline-alerts.rst +++ b/airflow-core/docs/howto/deadline-alerts.rst @@ -451,7 +451,8 @@ choose a different time. """A custom reference evaluated when Dag runs are queued.""" def _evaluate_with(self, *, session: Session, dagrun) -> datetime: - return dagrun.queued_at + my_datetime = my_business_logic(dagrun.queued_at) + return my_datetime **Using a Custom Reference in a Dag** diff --git a/airflow-core/src/airflow/models/deadline.py b/airflow-core/src/airflow/models/deadline.py index 332a8b6541142..d5e9aed564e94 100644 --- a/airflow-core/src/airflow/models/deadline.py +++ b/airflow-core/src/airflow/models/deadline.py @@ -51,6 +51,7 @@ from sqlalchemy.sql import ColumnElement from airflow.models.callback import CallbackDefinitionProtocol + from airflow.models.dagrun import DagRun from airflow.models.deadline_alert import DeadlineAlert @@ -357,7 +358,7 @@ def evaluate_with(self, *, session: Session, interval: timedelta, **kwargs: Any) return base_time + interval if base_time is not None else None @abstractmethod - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: """Must be implemented by subclasses to perform the actual evaluation.""" raise NotImplementedError @@ -407,13 +408,13 @@ def deserialize_reference(cls, reference_data: dict): class DagRunLogicalDateDeadline(BaseDeadlineReference): """A deadline that returns a DagRun's logical date.""" - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: return dagrun.logical_date class DagRunQueuedAtDeadline(BaseDeadlineReference): """A deadline that returns when a DagRun was queued.""" - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: return dagrun.queued_at @dataclass @@ -431,7 +432,7 @@ def __post_init__(self): raise ValueError("min_runs must be at least 1") @provide_session - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: from airflow.models import DagRun dag_id = dagrun.dag_id diff --git a/task-sdk/src/airflow/sdk/definitions/deadline.py b/task-sdk/src/airflow/sdk/definitions/deadline.py index 59be2eb6a248e..764a8f60173d8 100644 --- a/task-sdk/src/airflow/sdk/definitions/deadline.py +++ b/task-sdk/src/airflow/sdk/definitions/deadline.py @@ -218,7 +218,9 @@ class DeadlineReference: class MyDeadlineReference(BaseDeadlineReference): def _evaluate_with(self, *, session, dagrun): - return dagrun.logical_date + # Add your business logic here; dagrun.logical_date is available to use. + my_datetime = my_business_logic(dagrun.logical_date) + return my_datetime """ class TYPES: @@ -347,7 +349,9 @@ def _evaluate_with(self, *, session: Session, **kwargs) -> datetime: class MyCustomReference(BaseDeadlineReference): # By default, evaluate_with will be called when a new dagrun is created. def _evaluate_with(self, *, session: Session, dagrun) -> datetime: - return dagrun.logical_date + # Add your business logic here; dagrun.logical_date is available to use. + my_datetime = my_business_logic(dagrun.logical_date) + return my_datetime def serialize_reference(self) -> dict: return {"reference_type": self.reference_name} @@ -358,7 +362,9 @@ def serialize_reference(self) -> dict: class MyQueuedRef(BaseDeadlineReference): # Optionally, you can specify when you want it calculated by providing a DeadlineReference.TYPES def _evaluate_with(self, *, session: Session, dagrun) -> datetime: - return dagrun.queued_at + # Add your business logic here; dagrun.queued_at is available to use. + my_datetime = my_business_logic(dagrun.queued_at) + return my_datetime def serialize_reference(self) -> dict: return {"reference_type": self.reference_name} diff --git a/uv.lock b/uv.lock index 7a33696faabd9..fa137bb772ede 100644 --- a/uv.lock +++ b/uv.lock @@ -64,9 +64,9 @@ apache-airflow-providers-apache-cassandra = false apache-airflow-providers-asana = false apache-airflow-providers-oracle = false apache-airflow-providers-mysql = false +apache-airflow-providers-teradata = false apache-airflow-providers-alibaba = false apache-airflow-providers-microsoft-mssql = false -apache-airflow-providers-teradata = false apache-airflow-providers-jdbc = false apache-airflow-helm-chart = false apache-airflow-providers-anthropic = false @@ -4211,7 +4211,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-clickhousedb" -version = "1.0.1" +version = "1.0.0" source = { editable = "providers/clickhousedb" } dependencies = [ { name = "apache-airflow" }, @@ -6861,7 +6861,7 @@ docs = [{ name = "apache-airflow-devel-common", extras = ["docs"], editable = "d [[package]] name = "apache-airflow-providers-opensearch" -version = "1.11.2" +version = "1.12.0" source = { editable = "providers/opensearch" } dependencies = [ { name = "apache-airflow" }, From 6f6664ff4705af20da153fced1958b0ea709aef8 Mon Sep 17 00:00:00 2001 From: imrichardwu Date: Thu, 6 Aug 2026 22:55:14 -0700 Subject: [PATCH 6/8] Update deadline reference test for var-keyword context forwarding A custom deadline reference whose _evaluate_with declares **kwargs is expected to receive the full evaluation context, matching the standard Python "accept anything" convention rather than having the context filtered away. Align the test with that behaviour. --- .../tests/unit/models/test_deadline_alert.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/airflow-core/tests/unit/models/test_deadline_alert.py b/airflow-core/tests/unit/models/test_deadline_alert.py index 935f71beeb60e..48618eb3a9f08 100644 --- a/airflow-core/tests/unit/models/test_deadline_alert.py +++ b/airflow-core/tests/unit/models/test_deadline_alert.py @@ -243,25 +243,30 @@ def _evaluate_with(self, *, session, dagrun, **kwargs): run_id="test_run", ) == DEFAULT_DATE + timedelta(hours=1) - def test_legacy_serialized_custom_reference_ignores_evaluation_context(self): - class LegacyReference(BaseDeadlineReference): + def test_serialized_custom_reference_var_keyword_receives_full_context(self): + class VarKeywordReference(BaseDeadlineReference): received_kwargs: dict[str, object] def _evaluate_with(self, *, session, **kwargs): self.received_kwargs = kwargs return DEFAULT_DATE - reference = LegacyReference() + reference = VarKeywordReference() wrapper = SerializedReferenceModels.SerializedCustomReference(reference) + dagrun = SimpleNamespace() assert wrapper.evaluate_with( session=None, interval=timedelta(hours=1), - dagrun=SimpleNamespace(), + dagrun=dagrun, dag_id="test_dag", run_id="test_run", ) == DEFAULT_DATE + timedelta(hours=1) - assert reference.received_kwargs == {} + assert reference.received_kwargs == { + "dagrun": dagrun, + "dag_id": "test_dag", + "run_id": "test_run", + } def test_core_deadline_reference_treated_as_builtins(self): """Test that refs from airflow.models.deadline are still treated as builtins.""" From 6efefa1d951c95a49001075e5f9e7ae02faa73c6 Mon Sep 17 00:00:00 2001 From: imrichardwu Date: Mon, 10 Aug 2026 18:13:33 -0700 Subject: [PATCH 7/8] Improve deadline reference type safety Clarify the supported Dag run interface for deadline references and document the migration away from required_kwargs before Airflow 4.0. --- airflow-core/docs/howto/deadline-alerts.rst | 3 +++ airflow-core/src/airflow/models/deadline.py | 23 ++++++++++++++----- .../serialization/definitions/deadline.py | 11 +++++---- .../src/airflow/sdk/definitions/deadline.py | 1 - 4 files changed, 26 insertions(+), 12 deletions(-) diff --git a/airflow-core/docs/howto/deadline-alerts.rst b/airflow-core/docs/howto/deadline-alerts.rst index 5dc07759f92af..729c76768edd8 100644 --- a/airflow-core/docs/howto/deadline-alerts.rst +++ b/airflow-core/docs/howto/deadline-alerts.rst @@ -526,4 +526,7 @@ followed by a more urgent escalation if the Dag is still running. * **Plugin Placement**: One convenient place for custom references is in the plugins directory. * **API Server Restart**: Restart the Airflow API Server after adding or modifying custom references. * **Dag run context**: Add a ``dagrun`` parameter to ``_evaluate_with()`` to use the Dag run being evaluated. +* **Required Parameters** *(deprecated)*: ``required_kwargs`` declares which parameters your reference + needs. It still works but is deprecated and will be removed in Airflow 4.0 — declare the keyword-only + parameters your ``_evaluate_with()`` needs instead. * **Database Access**: Use the ``session`` parameter for Airflow database queries if needed. diff --git a/airflow-core/src/airflow/models/deadline.py b/airflow-core/src/airflow/models/deadline.py index d5e9aed564e94..afdb2d055c221 100644 --- a/airflow-core/src/airflow/models/deadline.py +++ b/airflow-core/src/airflow/models/deadline.py @@ -23,7 +23,7 @@ from dataclasses import dataclass from datetime import datetime, timedelta from inspect import signature -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, Protocol, cast from uuid import UUID import uuid6 @@ -40,6 +40,7 @@ ExecutorCallback, TriggererCallback, ) +from airflow.sdk.types import DagRunProtocol from airflow.utils.helpers import prune_dict from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.session import provide_session @@ -51,7 +52,6 @@ from sqlalchemy.sql import ColumnElement from airflow.models.callback import CallbackDefinitionProtocol - from airflow.models.dagrun import DagRun from airflow.models.deadline_alert import DeadlineAlert @@ -60,6 +60,17 @@ CALLBACK_METRICS_PREFIX = "deadline_alerts" +class DeadlineDagRunProtocol(DagRunProtocol, Protocol): + """ + The Dag run a deadline reference is evaluated against. + + Deadlines are only ever evaluated server-side against an ORM DagRun, so this exposes + ``queued_at`` on top of the fields a Dag run offers during task execution. + """ + + queued_at: datetime | None + + def _get_evaluation_kwargs(reference: Any, evaluator: Any, kwargs: dict[str, Any]) -> dict[str, Any]: """Return the evaluation arguments accepted by a deadline reference.""" required_kwargs: set[str] | None = getattr(reference, "required_kwargs", None) @@ -358,7 +369,7 @@ def evaluate_with(self, *, session: Session, interval: timedelta, **kwargs: Any) return base_time + interval if base_time is not None else None @abstractmethod - def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: """Must be implemented by subclasses to perform the actual evaluation.""" raise NotImplementedError @@ -408,13 +419,13 @@ def deserialize_reference(cls, reference_data: dict): class DagRunLogicalDateDeadline(BaseDeadlineReference): """A deadline that returns a DagRun's logical date.""" - def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: return dagrun.logical_date class DagRunQueuedAtDeadline(BaseDeadlineReference): """A deadline that returns when a DagRun was queued.""" - def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: return dagrun.queued_at @dataclass @@ -432,7 +443,7 @@ def __post_init__(self): raise ValueError("min_runs must be at least 1") @provide_session - def _evaluate_with(self, *, session: Session, dagrun: DagRun) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: from airflow.models import DagRun dag_id = dagrun.dag_id diff --git a/airflow-core/src/airflow/serialization/definitions/deadline.py b/airflow-core/src/airflow/serialization/definitions/deadline.py index f39112d831733..c3d8f22205568 100644 --- a/airflow-core/src/airflow/serialization/definitions/deadline.py +++ b/airflow-core/src/airflow/serialization/definitions/deadline.py @@ -38,6 +38,7 @@ from sqlalchemy import ColumnElement from sqlalchemy.orm import Session + from airflow.models.deadline import DeadlineDagRunProtocol from airflow.sdk.definitions.deadline import VariableInterval logger = logging.getLogger(__name__) @@ -107,7 +108,7 @@ def evaluate_with(self, *, session: Session, interval: timedelta, **kwargs: Any) return base_time + interval if base_time is not None else None @abstractmethod - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: """Must be implemented by subclasses to perform the actual evaluation.""" raise NotImplementedError @@ -138,7 +139,7 @@ class FixedDatetimeDeadline(SerializedBaseDeadlineReference): _datetime: datetime - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: return self._datetime def serialize_reference(self) -> dict: @@ -154,13 +155,13 @@ def deserialize_reference(cls, reference_data: dict): class DagRunLogicalDateDeadline(SerializedBaseDeadlineReference): """A deadline that returns a DagRun's logical date.""" - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: return dagrun.logical_date class DagRunQueuedAtDeadline(SerializedBaseDeadlineReference): """A deadline that returns when a DagRun was queued.""" - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: return dagrun.queued_at @dataclass @@ -178,7 +179,7 @@ def __post_init__(self): raise ValueError("min_runs must be at least 1") @provide_session - def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, dagrun: DeadlineDagRunProtocol) -> datetime | None: from sqlalchemy import func, text from airflow.models import DagRun diff --git a/task-sdk/src/airflow/sdk/definitions/deadline.py b/task-sdk/src/airflow/sdk/definitions/deadline.py index 764a8f60173d8..0a2cdf988055f 100644 --- a/task-sdk/src/airflow/sdk/definitions/deadline.py +++ b/task-sdk/src/airflow/sdk/definitions/deadline.py @@ -360,7 +360,6 @@ def serialize_reference(self) -> dict: # Optionally, specify when it is calculated by providing a DeadlineReference.TYPES value. @deadline_reference(DeadlineReference.TYPES.DAGRUN_QUEUED) class MyQueuedRef(BaseDeadlineReference): - # Optionally, you can specify when you want it calculated by providing a DeadlineReference.TYPES def _evaluate_with(self, *, session: Session, dagrun) -> datetime: # Add your business logic here; dagrun.queued_at is available to use. my_datetime = my_business_logic(dagrun.queued_at) From f4806fb0f1dab3f387257a9936b425ed2b65d47c Mon Sep 17 00:00:00 2001 From: imrichardwu Date: Thu, 13 Aug 2026 23:33:02 -0700 Subject: [PATCH 8/8] Refactor DeadlineDagRunProtocol to remove dependency on DagRunProtocol Updated the DeadlineDagRunProtocol class to eliminate its inheritance from DagRunProtocol, simplifying the interface for deadline references. The class now directly defines the necessary attributes for evaluating deadlines without relying on the ORM DagRun structure. --- airflow-core/src/airflow/models/deadline.py | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/airflow-core/src/airflow/models/deadline.py b/airflow-core/src/airflow/models/deadline.py index afdb2d055c221..ed817bf07a736 100644 --- a/airflow-core/src/airflow/models/deadline.py +++ b/airflow-core/src/airflow/models/deadline.py @@ -40,7 +40,6 @@ ExecutorCallback, TriggererCallback, ) -from airflow.sdk.types import DagRunProtocol from airflow.utils.helpers import prune_dict from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.session import provide_session @@ -60,14 +59,11 @@ CALLBACK_METRICS_PREFIX = "deadline_alerts" -class DeadlineDagRunProtocol(DagRunProtocol, Protocol): - """ - The Dag run a deadline reference is evaluated against. - - Deadlines are only ever evaluated server-side against an ORM DagRun, so this exposes - ``queued_at`` on top of the fields a Dag run offers during task execution. - """ +class DeadlineDagRunProtocol(Protocol): + """The Dag run interface deadline references are evaluated against.""" + dag_id: str + logical_date: datetime | None queued_at: datetime | None