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
Original file line number Diff line number Diff line change
Expand Up @@ -63,30 +63,20 @@ class DateTimeSensor(BaseSensorOperator):

def __init__(self, *, target_time: str | datetime.datetime, **kwargs) -> None:
super().__init__(**kwargs)

# self.target_time can't be a datetime object as it is a template_field
if isinstance(target_time, datetime.datetime):
self.target_time = target_time.isoformat()
Comment thread
shahar1 marked this conversation as resolved.
elif isinstance(target_time, str):
self.target_time = target_time
else:
raise TypeError(
f"Expected str or datetime.datetime type for target_time. Got {type(target_time)}"
)
self.target_time = target_time

def poke(self, context: Context) -> bool:
self.log.info("Checking if the time (%s) has come", self.target_time)
return timezone.utcnow() > timezone.parse(self.target_time)
return timezone.utcnow() > self._moment

@property
def _moment(self) -> datetime.datetime:
# Note following is reachable code if Jinja is used for redering template fields and
# render_template_as_native_obj=True is used.
# In this case, the target_time is already a datetime object.
if isinstance(self.target_time, datetime.datetime): # type:ignore[unreachable]
return self.target_time # type:ignore[unreachable]

return timezone.parse(self.target_time)
target_time: Any = self.target_time
if isinstance(target_time, datetime.datetime):
target_time = target_time.isoformat()
if isinstance(target_time, str):
return timezone.parse(target_time)
raise TypeError(f"Expected str or datetime.datetime type for target_time. Got {type(target_time)}")


class DateTimeSensorAsync(DateTimeSensor):
Expand Down Expand Up @@ -126,7 +116,7 @@ def __init__(
self.start_from_trigger = start_from_trigger
if self.start_from_trigger:
self.start_trigger_args.trigger_kwargs = dict(
moment=timezone.parse(self.target_time),
moment=self._moment,
end_from_trigger=self.end_from_trigger,
)

Expand Down
85 changes: 59 additions & 26 deletions providers/standard/tests/unit/standard/sensors/test_date_time.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,14 +17,15 @@
# under the License.
from __future__ import annotations

import datetime
from unittest.mock import patch

import pendulum
import pytest

from airflow import macros
from airflow.models.dag import DAG
from airflow.providers.standard.sensors.date_time import DateTimeSensor
from airflow.providers.standard.sensors.date_time import DateTimeSensor, DateTimeSensorAsync

from tests_common.test_utils.version_compat import timezone

Expand All @@ -38,41 +39,30 @@ def setup_class(cls):
cls.dag = DAG("test_dag", schedule=None, default_args=args)

@pytest.mark.parametrize(
("task_id", "target_time", "expected"),
("task_id", "target_time"),
[
(
"valid_datetime",
timezone.datetime(2020, 7, 6, 13, tzinfo=timezone.utc),
"2020-07-06T13:00:00+00:00",
),
(
"valid_str",
"20200706T210000+8",
"20200706T210000+8",
),
(
"jinja_str_is_accepted",
"{{ ds }}",
"{{ ds }}",
),
("valid_datetime", timezone.datetime(2020, 7, 6, 13, tzinfo=timezone.utc)),
("valid_str", "20200706T210000+8"),
("jinja_str_is_accepted", "{{ ds }}"),
],
)
def test_valid_input(self, task_id, target_time, expected):
"""target_time should be a string as it is a template field"""
def test_target_time_stored_verbatim(self, task_id, target_time):
"""target_time is a template field, so __init__ must store it as-is without transformation."""
op = DateTimeSensor(
task_id=task_id,
target_time=target_time,
dag=self.dag,
)
assert op.target_time == expected
assert op.target_time == target_time

def test_invalid_input(self):
def test_invalid_input_rejected_after_rendering(self):
op = DateTimeSensor(
task_id="test",
target_time=timezone.utcnow().time(),
dag=self.dag,
)
with pytest.raises(TypeError):
DateTimeSensor(
task_id="test",
target_time=timezone.utcnow().time(),
dag=self.dag,
)
op.poke(None)

@pytest.mark.parametrize(
("task_id", "target_time", "expected"),
Expand Down Expand Up @@ -124,3 +114,46 @@ def test_moment(self, native, target_time, expected_type):
sensor.render_template_fields(ctx)

assert isinstance(sensor._moment, expected_type)

@patch(
"airflow.providers.standard.sensors.date_time.timezone.utcnow",
return_value=pendulum.datetime(2020, 1, 2, tz="UTC"),
)
def test_poke_with_natively_rendered_datetime(self, mock_utcnow):
"""poke handles a target_time that native rendering resolved to a datetime."""
dag = DAG(
dag_id="native_poke_dag",
start_date=pendulum.datetime(2025, 1, 1, tz="UTC"),
schedule=None,
render_template_as_native_obj=True,
)
op = DateTimeSensor(task_id="native_poke", target_time="{{ data_interval_end }}", dag=dag)
op.render_template_fields(
{"data_interval_end": pendulum.datetime(2020, 1, 1, tz="UTC"), "macros": macros, "dag": dag}
)
assert isinstance(op.target_time, datetime.datetime)
assert op.poke(None) is True

def test_moment_localizes_naive_datetime(self):
"""A naive datetime target_time is localized to UTC via _moment (mirrors old isoformat())."""
op = DateTimeSensor(task_id="naive", target_time=datetime.datetime(2020, 1, 1), dag=self.dag)
assert op._moment == pendulum.datetime(2020, 1, 1, tz="UTC")

def test_async_start_from_trigger_moment(self):
op = DateTimeSensorAsync(
task_id="async",
target_time="2020-01-01T00:00:00+00:00",
start_from_trigger=True,
dag=self.dag,
)
assert op.start_trigger_args.trigger_kwargs["moment"] == pendulum.parse("2020-01-01T00:00:00+00:00")

def test_async_start_from_trigger_localizes_naive_datetime(self):
"""DateTimeSensorAsync never pokes, so _moment must still localize a naive datetime."""
op = DateTimeSensorAsync(
task_id="async_naive",
target_time=datetime.datetime(2020, 1, 1),
start_from_trigger=True,
dag=self.dag,
)
assert op.start_trigger_args.trigger_kwargs["moment"] == pendulum.datetime(2020, 1, 1, tz="UTC")
1 change: 0 additions & 1 deletion scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -68,4 +68,3 @@ providers/ssh/src/airflow/providers/ssh/operators/ssh.py::SSHOperator
providers/ssh/src/airflow/providers/ssh/operators/ssh_remote_job.py::SSHRemoteJobOperator
providers/standard/src/airflow/providers/standard/operators/bash.py::BashOperator
providers/standard/src/airflow/providers/standard/operators/trigger_dagrun.py::TriggerDagRunOperator
providers/standard/src/airflow/providers/standard/sensors/date_time.py::DateTimeSensor
Loading