diff --git a/airflow-core/docs/howto/deadline-alerts.rst b/airflow-core/docs/howto/deadline-alerts.rst index 124b7d87040c5..7a85c3dab8d94 100644 --- a/airflow-core/docs/howto/deadline-alerts.rst +++ b/airflow-core/docs/howto/deadline-alerts.rst @@ -451,9 +451,9 @@ Place the reference classes and the plugin that registers them in your plugins f 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: + 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. @@ -461,14 +461,9 @@ Place the reference classes and the plugin that registers them in your plugins f 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: + my_datetime = my_business_logic(dagrun.queued_at) + return my_datetime # Register the classes so the scheduler can resolve them when it deserializes the Dag. @@ -549,8 +544,8 @@ followed by a more urgent escalation if the Dag is still running. * **Plugin Registration**: Custom references must be listed in the ``deadline_references`` attribute of an ``AirflowPlugin``, so the plugins directory is the natural home for them. * **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. +* **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/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..ed817bf07a736 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 typing import TYPE_CHECKING, Any, cast +from inspect import signature +from typing import TYPE_CHECKING, Any, Protocol, 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,38 @@ CALLBACK_METRICS_PREFIX = "deadline_alerts" +class DeadlineDagRunProtocol(Protocol): + """The Dag run interface deadline references are evaluated against.""" + + dag_id: str + logical_date: datetime | None + 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) + 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 + 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} + + class classproperty: """ Decorator that converts a method with a single cls argument into a property. @@ -320,34 +354,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: DeadlineDagRunProtocol) -> datetime | None: """Must be implemented by subclasses to perform the actual evaluation.""" raise NotImplementedError @@ -381,7 +399,7 @@ class FixedDatetimeDeadline(BaseDeadlineReference): _datetime: datetime - def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None: + def _evaluate_with(self, *, session: Session, **kwargs) -> datetime | None: return self._datetime def serialize_reference(self) -> dict: @@ -397,23 +415,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: DeadlineDagRunProtocol) -> 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: DeadlineDagRunProtocol) -> datetime | None: + return dagrun.queued_at @dataclass class AverageRuntimeDeadline(BaseDeadlineReference): @@ -422,7 +431,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 +439,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: DeadlineDagRunProtocol) -> 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 +515,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 6a91d18223e3c..34c1356bb5b72 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..c3d8f22205568 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 @@ -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__) @@ -96,33 +97,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: DeadlineDagRunProtocol) -> datetime | None: """Must be implemented by subclasses to perform the actual evaluation.""" raise NotImplementedError @@ -169,23 +155,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: DeadlineDagRunProtocol) -> 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: DeadlineDagRunProtocol) -> datetime | None: + return dagrun.queued_at @dataclass class AverageRuntimeDeadline(SerializedBaseDeadlineReference): @@ -194,7 +171,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 +179,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: DeadlineDagRunProtocol) -> 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 +256,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 +267,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 +320,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 b52bbd5120b99..356143d0a3c0d 100644 --- a/airflow-core/tests/unit/models/test_deadline.py +++ b/airflow-core/tests/unit/models/test_deadline.py @@ -19,17 +19,17 @@ import re from dataclasses import dataclass 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 @@ -55,15 +55,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.""" @@ -318,138 +309,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.""" @@ -480,7 +374,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) @@ -517,7 +413,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 @@ -551,7 +449,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 @@ -564,7 +464,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): @@ -613,7 +515,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 @@ -637,61 +541,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..48618eb3a9f08 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,83 @@ 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_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 = VarKeywordReference() + wrapper = SerializedReferenceModels.SerializedCustomReference(reference) + dagrun = SimpleNamespace() + + 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) + 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.""" diff --git a/airflow-core/tests/unit/models/test_taskinstance.py b/airflow-core/tests/unit/models/test_taskinstance.py index 6a348a3f2ca26..035c21280ac25 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): +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), @@ -4188,6 +4208,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 f3a14aeb9cc55..416324c68ef83 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). A custom + ``_evaluate_with()`` with deferred Core imports (imports inside the method body). Its + keyword-only ``dagrun`` parameter receives the DagRun being evaluated. A custom reference must be decorated with ``@deadline_reference`` and listed in the ``deadline_references`` attribute of an ``AirflowPlugin``; see :external:doc:`howto/deadline-alerts`. """ @@ -185,8 +186,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: @@ -213,17 +214,15 @@ 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): + # 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: @@ -368,11 +367,10 @@ def _evaluate_with(self, *, session: Session, **kwargs) -> datetime: @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: + # 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} @@ -381,9 +379,10 @@ 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): - def _evaluate_with(self, *, session: Session, **kwargs) -> datetime: - # Put your business logic here - return some_datetime + 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) + return my_datetime def serialize_reference(self) -> dict: return {"reference_type": self.reference_name}