Skip to content

Commit db15212

Browse files
committed
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.
1 parent 9d07d7c commit db15212

10 files changed

Lines changed: 221 additions & 398 deletions

File tree

‎airflow-core/docs/howto/deadline-alerts.rst‎

Lines changed: 5 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -437,24 +437,17 @@ implement an ``_evaluate_with()`` method.
437437
class MyCustomDecoratedReference(BaseDeadlineReference):
438438
"""A custom reference evaluated when Dag runs are created."""
439439
440-
def _evaluate_with(self, *, session: Session, **kwargs) -> datetime:
441-
# Add your business logic here
442-
return your_datetime
440+
def _evaluate_with(self, *, session: Session, dagrun) -> datetime:
441+
return dagrun.logical_date
443442
444443
445444
# You can specify when evaluate_with will be called by providing a DeadlineReference.TYPES value.
446445
@deadline_reference(DeadlineReference.TYPES.DAGRUN_QUEUED)
447446
class MyQueuedReference(BaseDeadlineReference):
448447
"""A custom reference evaluated when Dag runs are queued."""
449448
450-
# Ask for the Dag run context values supplied by Airflow; see notes below.
451-
required_kwargs = {"dag_id", "run_id"}
452-
453-
def _evaluate_with(self, *, session: Session, **kwargs) -> datetime:
454-
dag_id = kwargs["dag_id"]
455-
run_id = kwargs["run_id"]
456-
# Use dag_id and run_id in your calculation
457-
return your_datetime
449+
def _evaluate_with(self, *, session: Session, dagrun) -> datetime:
450+
return dagrun.queued_at
458451
459452
460453
**Using a Custom Reference in a Dag**
@@ -524,8 +517,5 @@ followed by a more urgent escalation if the Dag is still running.
524517
* **Timezone Awareness**: Always return timezone-aware datetime objects.
525518
* **Plugin Placement**: One convenient place for custom references is in the plugins directory.
526519
* **API Server Restart**: Restart the Airflow API Server after adding or modifying custom references.
527-
* **Required Parameters**: ``required_kwargs`` declares which Dag run context values Airflow should
528-
forward to ``_evaluate_with()``. Only ``dag_id`` and ``run_id`` are available; declaring anything
529-
else raises a ``ValueError`` when the deadline is evaluated. To configure a reference itself, give
530-
it constructor fields or read from an Airflow Variable.
520+
* **Dag run context**: Add a ``dagrun`` parameter to ``_evaluate_with()`` to use the Dag run being evaluated.
531521
* **Database Access**: Use the ``session`` parameter for Airflow database queries if needed.
Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
Deprecate the ``required_kwargs`` attribute for custom deadline references. Implement ``_evaluate_with()`` with a ``dagrun`` parameter instead.

‎airflow-core/src/airflow/models/deadline.py‎

Lines changed: 36 additions & 90 deletions
Original file line numberDiff line numberDiff line change
@@ -17,21 +17,23 @@
1717
from __future__ import annotations
1818

1919
import logging
20+
import warnings
2021
from abc import ABC, abstractmethod
2122
from collections.abc import Sequence
2223
from dataclasses import dataclass
2324
from datetime import datetime, timedelta
25+
from inspect import signature
2426
from typing import TYPE_CHECKING, Any, cast
2527
from uuid import UUID
2628

2729
import uuid6
28-
from sqlalchemy import Boolean, ForeignKey, Index, Integer, Uuid, and_, func, inspect, select, text
29-
from sqlalchemy.exc import SQLAlchemyError
30+
from sqlalchemy import Boolean, ForeignKey, Index, Integer, Uuid, and_, func, select, text
3031
from sqlalchemy.orm import Mapped, mapped_column, relationship
3132

3233
from airflow._shared.observability.metrics import stats
3334
from airflow._shared.timezones import timezone
3435
from airflow.configuration import conf
36+
from airflow.exceptions import RemovedInAirflow4Warning
3537
from airflow.models.base import Base
3638
from airflow.models.callback import (
3739
Callback,
@@ -57,6 +59,27 @@
5759
CALLBACK_METRICS_PREFIX = "deadline_alerts"
5860

5961

62+
def _get_evaluation_kwargs(reference: Any, evaluator: Any, kwargs: dict[str, Any]) -> dict[str, Any]:
63+
"""Return the evaluation arguments accepted by a deadline reference."""
64+
required_kwargs: set[str] | None = getattr(reference, "required_kwargs", None)
65+
if required_kwargs is not None:
66+
warnings.warn(
67+
"required_kwargs is deprecated. Declare the keyword-only parameters your "
68+
"_evaluate_with() implementation needs instead.",
69+
RemovedInAirflow4Warning,
70+
stacklevel=3,
71+
)
72+
kwargs = {key: value for key, value in kwargs.items() if key in required_kwargs}
73+
if missing_kwargs := required_kwargs - kwargs.keys():
74+
raise ValueError(
75+
f"{reference.__class__.__name__} is missing required parameters: {', '.join(missing_kwargs)}"
76+
)
77+
return kwargs
78+
79+
parameters = signature(evaluator).parameters
80+
return {key: value for key, value in kwargs.items() if key in parameters}
81+
82+
6083
class classproperty:
6184
"""
6285
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]
320343
class BaseDeadlineReference(LoggingMixin, ABC):
321344
"""Base class for all Deadline implementations."""
322345

323-
# Set of required kwargs - subclasses should override this.
324-
required_kwargs: set[str] = set()
325-
326346
@classproperty
327347
def reference_name(cls: Any) -> str:
328348
return cls.__name__
329349

330350
def evaluate_with(self, *, session: Session, interval: timedelta, **kwargs: Any) -> datetime | None:
331-
"""Validate the provided kwargs and evaluate this deadline with the given conditions."""
332-
filtered_kwargs = {k: v for k, v in kwargs.items() if k in self.required_kwargs}
333-
334-
if missing_kwargs := self.required_kwargs - filtered_kwargs.keys():
335-
raise ValueError(
336-
f"{self.__class__.__name__} is missing required parameters: {', '.join(missing_kwargs)}"
337-
)
338-
339-
if extra_kwargs := kwargs.keys() - filtered_kwargs.keys():
340-
self.log.debug(
341-
"%s ignoring unexpected parameters: %s",
342-
self.reference_name,
343-
", ".join(extra_kwargs),
344-
)
345-
346-
base_time = self._evaluate_with(session=session, **filtered_kwargs)
351+
"""Evaluate this deadline with the supplied context."""
352+
evaluation_kwargs = _get_evaluation_kwargs(self, self._evaluate_with, kwargs)
353+
base_time = self._evaluate_with(session=session, **evaluation_kwargs)
347354
return base_time + interval if base_time is not None else None
348355

349356
@abstractmethod
350-
def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None:
357+
def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None:
351358
"""Must be implemented by subclasses to perform the actual evaluation."""
352359
raise NotImplementedError
353360

@@ -381,7 +388,7 @@ class FixedDatetimeDeadline(BaseDeadlineReference):
381388

382389
_datetime: datetime
383390

384-
def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None:
391+
def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None:
385392
return self._datetime
386393

387394
def serialize_reference(self) -> dict:
@@ -397,23 +404,14 @@ def deserialize_reference(cls, reference_data: dict):
397404
class DagRunLogicalDateDeadline(BaseDeadlineReference):
398405
"""A deadline that returns a DagRun's logical date."""
399406

400-
required_kwargs = {"dag_id", "run_id"}
401-
402-
def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None:
403-
from airflow.models import DagRun
404-
405-
return _fetch_from_db(DagRun.logical_date, session=session, **kwargs)
407+
def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None:
408+
return dagrun.logical_date
406409

407410
class DagRunQueuedAtDeadline(BaseDeadlineReference):
408411
"""A deadline that returns when a DagRun was queued."""
409412

410-
required_kwargs = {"dag_id", "run_id"}
411-
412-
@provide_session
413-
def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None:
414-
from airflow.models import DagRun
415-
416-
return _fetch_from_db(DagRun.queued_at, session=session, **kwargs)
413+
def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None:
414+
return dagrun.queued_at
417415

418416
@dataclass
419417
class AverageRuntimeDeadline(BaseDeadlineReference):
@@ -422,7 +420,6 @@ class AverageRuntimeDeadline(BaseDeadlineReference):
422420
DEFAULT_LIMIT = 10
423421
max_runs: int
424422
min_runs: int | None = None
425-
required_kwargs = {"dag_id"}
426423

427424
def __post_init__(self):
428425
if self.min_runs is None:
@@ -431,10 +428,10 @@ def __post_init__(self):
431428
raise ValueError("min_runs must be at least 1")
432429

433430
@provide_session
434-
def _evaluate_with(self, *, session: Session, **kwargs: Any) -> datetime | None:
431+
def _evaluate_with(self, *, session: Session, dagrun: Any) -> datetime | None:
435432
from airflow.models import DagRun
436433

437-
dag_id = kwargs["dag_id"]
434+
dag_id = dagrun.dag_id
438435

439436
# Get database dialect to use appropriate time difference calculation
440437
dialect = get_dialect_name(session)
@@ -507,54 +504,3 @@ def deserialize_reference(cls, reference_data: dict):
507504

508505

509506
DeadlineReferenceType = ReferenceModels.BaseDeadlineReference
510-
511-
512-
@provide_session
513-
def _fetch_from_db(model_reference: Mapped, *, session=None, **conditions) -> datetime | None:
514-
"""
515-
Fetch a datetime value from the database using the provided model reference and filtering conditions.
516-
517-
For example, to fetch a TaskInstance's start_date:
518-
_fetch_from_db(
519-
TaskInstance.start_date, dag_id='example_dag', task_id='example_task', run_id='example_run'
520-
)
521-
522-
This generates SQL equivalent to:
523-
SELECT start_date
524-
FROM task_instance
525-
WHERE dag_id = 'example_dag'
526-
AND task_id = 'example_task'
527-
AND run_id = 'example_run'
528-
529-
:param model_reference: SQLAlchemy Column to select (e.g., DagRun.logical_date, TaskInstance.start_date)
530-
:param conditions: Filtering conditions applied as equality comparisons in the WHERE clause.
531-
Multiple conditions are combined with AND.
532-
:param session: SQLAlchemy session (auto-provided by decorator)
533-
"""
534-
query = select(model_reference)
535-
536-
for key, value in conditions.items():
537-
inspected = inspect(model_reference)
538-
if inspected is not None:
539-
query = query.where(getattr(inspected.class_, key) == value)
540-
541-
compiled_query = query.compile(compile_kwargs={"literal_binds": True})
542-
pretty_query = "\n ".join(str(compiled_query).splitlines())
543-
logger.debug(
544-
"Executing query:\n %r\nAs SQL:\n %s",
545-
query,
546-
pretty_query,
547-
)
548-
549-
try:
550-
result = session.scalar(query)
551-
except SQLAlchemyError:
552-
logger.exception("Database query failed.")
553-
raise
554-
555-
if result is None:
556-
message = f"No matching record found in the database for query:\n {pretty_query}"
557-
logger.error(message)
558-
raise ValueError(message)
559-
560-
return result

‎airflow-core/src/airflow/models/taskinstance.py‎

Lines changed: 13 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -94,6 +94,7 @@
9494
from airflow.models.taskmap import TaskMap
9595
from airflow.models.taskreschedule import TaskReschedule
9696
from airflow.models.xcom import XCOM_RETURN_KEY, LazyXComSelectSequence, XComModel
97+
from airflow.serialization.decoders import decode_deadline_reference
9798
from airflow.serialization.enums import stringify_encoding_keys
9899
from airflow.settings import task_instance_mutation_hook
99100
from airflow.task.priority_strategy import validate_and_load_priority_weight_strategy
@@ -222,14 +223,11 @@ def _add_and_prime_mapped_ti(
222223
set_committed_value(ti, "dag_run", dag_run)
223224

224225

225-
def _recalculate_dagrun_queued_at_deadlines(
226-
dagrun: DagRun, new_queued_at: datetime, session: Session
227-
) -> None:
226+
def _recalculate_dagrun_queued_at_deadlines(dagrun: DagRun, *, session: Session) -> None:
228227
"""
229228
Recalculate deadline times for deadlines that reference dagrun.queued_at.
230229
231230
:param dagrun: The DagRun whose deadlines should be recalculated
232-
:param new_queued_at: The new queued_at timestamp to use for calculation
233231
:param session: Database session
234232
235233
:meta private:
@@ -249,9 +247,16 @@ def _recalculate_dagrun_queued_at_deadlines(
249247
return
250248

251249
for deadline, deadline_alert in results:
252-
# We can't use evaluate_with() since the new queued_at is not written to the DB yet.
253-
deadline_interval = timedelta(seconds=deadline_alert.interval)
254-
new_deadline_time = new_queued_at + deadline_interval
250+
new_deadline_time = decode_deadline_reference(deadline_alert.reference).evaluate_with(
251+
session=session,
252+
interval=timedelta(seconds=deadline_alert.interval),
253+
dagrun=dagrun,
254+
dag_id=dagrun.dag_id,
255+
run_id=dagrun.run_id,
256+
)
257+
258+
if new_deadline_time is None:
259+
continue
255260

256261
log.debug(
257262
"Recalculating deadline %s for DagRun %s.%s: old=%s, new=%s",
@@ -459,7 +464,7 @@ def clear_task_instances(
459464
parent_context=parent_trace_context(dr.conf),
460465
)
461466

462-
_recalculate_dagrun_queued_at_deadlines(dr, dr.queued_at, session)
467+
_recalculate_dagrun_queued_at_deadlines(dr, session=session)
463468

464469
if dr.state in State.finished_dr_states:
465470
dr.state = dag_run_state

‎airflow-core/src/airflow/serialization/definitions/dag.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -761,7 +761,7 @@ def _process_dagrun_deadline_alerts(
761761
deadline_time = deserialized_deadline_alert.reference.evaluate_with(
762762
session=session,
763763
interval=interval,
764-
# TODO : Pretty sure we can drop these last two; verify after testing is complete
764+
dagrun=orm_dagrun,
765765
dag_id=self.dag_id,
766766
run_id=orm_dagrun.run_id,
767767
)

0 commit comments

Comments
 (0)