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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 10 additions & 15 deletions airflow-core/docs/howto/deadline-alerts.rst
Original file line number Diff line number Diff line change
Expand Up @@ -451,24 +451,19 @@ 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.
@deadline_reference(DeadlineReference.TYPES.DAGRUN_QUEUED)
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.
Expand Down Expand Up @@ -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
Comment thread
imrichardwu marked this conversation as resolved.
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.
1 change: 1 addition & 0 deletions airflow-core/newsfragments/70714.significant.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Deprecate the ``required_kwargs`` attribute for custom deadline references. Implement ``_evaluate_with()`` with a ``dagrun`` parameter instead.
139 changes: 48 additions & 91 deletions airflow-core/src/airflow/models/deadline.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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}
Comment thread
imrichardwu marked this conversation as resolved.


class classproperty:
"""
Decorator that converts a method with a single cls argument into a property.
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand All @@ -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):
Expand All @@ -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:
Expand All @@ -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)
Expand Down Expand Up @@ -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
21 changes: 13 additions & 8 deletions airflow-core/src/airflow/models/taskinstance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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",
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
Loading