diff --git a/providers/standard/src/airflow/providers/standard/sensors/time_delta.py b/providers/standard/src/airflow/providers/standard/sensors/time_delta.py index 0f9bcf45d0767..09896c8779a76 100644 --- a/providers/standard/src/airflow/providers/standard/sensors/time_delta.py +++ b/providers/standard/src/airflow/providers/standard/sensors/time_delta.py @@ -57,10 +57,12 @@ class TimeDeltaSensor(BaseSensorOperator): """ + template_fields: Sequence[str] = ("delta",) + def __init__( self, *, - delta: timedelta, + delta: timedelta | int, deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), end_from_trigger: bool = False, **kwargs, @@ -70,6 +72,12 @@ def __init__( self.deferrable = deferrable self.end_from_trigger = end_from_trigger + def _resolve_delta(self) -> timedelta: + value = self.delta + if isinstance(value, timedelta): + return value + return timedelta(minutes=int(value)) + def _derive_base_time(self, context: Context) -> datetime: """ Get the "base time" against which the delta should be calculated. @@ -93,8 +101,9 @@ def _derive_base_time(self, context: Context) -> datetime: def poke(self, context: Context) -> bool: base_time = self._derive_base_time(context=context) - target_dttm = base_time + self.delta - self.log.info("Checking if the delta has elapsed base_time=%s, delta=%s", base_time, self.delta) + delta = self._resolve_delta() + target_dttm = base_time + delta + self.log.info("Checking if the delta has elapsed base_time=%s, delta=%s", base_time, delta) return timezone.utcnow() > target_dttm """ @@ -113,7 +122,7 @@ def execute(self, context: Context) -> Any: # Deferrable path base_time = self._derive_base_time(context=context) - target_dttm: datetime = base_time + self.delta + target_dttm: datetime = base_time + self._resolve_delta() if timezone.utcnow() > target_dttm: # If the target datetime is in the past, return immediately diff --git a/providers/standard/tests/unit/standard/sensors/test_time_delta.py b/providers/standard/tests/unit/standard/sensors/test_time_delta.py index 9ec8fc4233d5b..31b104f703cb5 100644 --- a/providers/standard/tests/unit/standard/sensors/test_time_delta.py +++ b/providers/standard/tests/unit/standard/sensors/test_time_delta.py @@ -255,6 +255,27 @@ def test_timedelta_sensor_async_run_after_vs_interval(self, run_after, interval_ assert caught.value.trigger.moment == expected_time + @pytest.mark.parametrize( + ("delta", "should_defer"), [(timedelta(minutes=1), True), (1, False), ("{{ 1*2 }}", True)] + ) + def test_time_delta_sensor_templating(self, mocker, delta, should_defer): + defer_mock = mocker.patch(DEFER_PATH) + op = TimeDeltaSensor(task_id="time_sensor_check", delta=delta, dag=self.dag, deferrable=should_defer) + + time = pendulum.datetime(year=2024, month=8, day=1, tz="UTC") + + with time_machine.travel(time, tick=False): + if should_defer: + data_interval_end = time.add(hours=1) + else: + data_interval_end = time.subtract(minutes=2) + + context = {"data_interval_end": data_interval_end} + op.render_template_fields(context) + op.execute(context) + if should_defer: + defer_mock.assert_called_once() + @pytest.mark.parametrize( "time_to_wait", [timedelta(minutes=1), 1, "{{ 1*2 }}"],