1717from __future__ import annotations
1818
1919import logging
20+ import warnings
2021from abc import ABC , abstractmethod
2122from collections .abc import Sequence
2223from dataclasses import dataclass
2324from datetime import datetime , timedelta
25+ from inspect import signature
2426from typing import TYPE_CHECKING , Any , cast
2527from uuid import UUID
2628
2729import 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
3031from sqlalchemy .orm import Mapped , mapped_column , relationship
3132
3233from airflow ._shared .observability .metrics import stats
3334from airflow ._shared .timezones import timezone
3435from airflow .configuration import conf
36+ from airflow .exceptions import RemovedInAirflow4Warning
3537from airflow .models .base import Base
3638from airflow .models .callback import (
3739 Callback ,
5759CALLBACK_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+
6083class 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
509506DeadlineReferenceType = 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\n As 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
0 commit comments