Skip to content
Merged
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 @@ -17,6 +17,7 @@
# under the License.
from __future__ import annotations

import dataclasses
import datetime
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, NoReturn
Expand Down Expand Up @@ -115,9 +116,15 @@ def __init__(

self.start_from_trigger = start_from_trigger
if self.start_from_trigger:
self.start_trigger_args.trigger_kwargs = dict(
moment=self._moment,
end_from_trigger=self.end_from_trigger,
# Replaced rather than mutated: ``start_trigger_args`` is a class attribute, so
# assigning through it would overwrite the arguments of every other task built
# from this operator.
self.start_trigger_args = dataclasses.replace(
self.start_trigger_args,
trigger_kwargs=dict(
moment=self._moment,
end_from_trigger=self.end_from_trigger,
),
)

def execute(self, context: Context) -> NoReturn:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
# under the License.
from __future__ import annotations

import dataclasses
import datetime
import os
from collections.abc import Sequence
Expand Down Expand Up @@ -90,11 +91,17 @@ def __init__(
self.start_from_trigger = start_from_trigger

if self.deferrable and self.start_from_trigger:
self.start_trigger_args.timeout = datetime.timedelta(seconds=self.timeout)
self.start_trigger_args.trigger_kwargs = dict(
filepath=self.path,
recursive=self.recursive,
poke_interval=self.poke_interval,
# Replaced rather than mutated: ``start_trigger_args`` is a class attribute, so
# assigning through it would overwrite the arguments of every other task built
# from this operator.
self.start_trigger_args = dataclasses.replace(
self.start_trigger_args,
timeout=datetime.timedelta(seconds=self.timeout),
trigger_kwargs=dict(
filepath=self.path,
recursive=self.recursive,
poke_interval=self.poke_interval,
),
)

@cached_property
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
# under the License.
from __future__ import annotations

import dataclasses
import datetime
import warnings
from typing import TYPE_CHECKING, Any
Expand Down Expand Up @@ -81,8 +82,12 @@ def __init__(
self.end_from_trigger = end_from_trigger

if self.start_from_trigger:
self.start_trigger_args.trigger_kwargs = dict(
moment=self.target_datetime, end_from_trigger=self.end_from_trigger
# Replaced rather than mutated: ``start_trigger_args`` is a class attribute, so
# assigning through it would overwrite the arguments of every other task built
# from this operator.
self.start_trigger_args = dataclasses.replace(
self.start_trigger_args,
trigger_kwargs=dict(moment=self.target_datetime, end_from_trigger=self.end_from_trigger),
)

def execute(self, context: Context) -> None:
Expand Down
32 changes: 32 additions & 0 deletions providers/standard/tests/unit/standard/sensors/test_date_time.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,3 +157,35 @@ def test_async_start_from_trigger_localizes_naive_datetime(self):
dag=self.dag,
)
assert op.start_trigger_args.trigger_kwargs["moment"] == pendulum.datetime(2020, 1, 1, tz="UTC")

def test_start_trigger_args_are_not_shared_between_tasks(self):
"""Each task must carry its own trigger arguments.

``start_trigger_args`` is a class attribute, so assigning through it made every task
built from this operator advertise the moment of whichever was constructed last.
"""
first = DateTimeSensorAsync(
task_id="first",
target_time="2030-01-01T00:00:00+00:00",
start_from_trigger=True,
dag=self.dag,
)
second = DateTimeSensorAsync(
task_id="second",
target_time="2040-06-06T00:00:00+00:00",
start_from_trigger=True,
dag=self.dag,
)

assert first.start_trigger_args is not second.start_trigger_args
assert first.start_trigger_args.trigger_kwargs["moment"] == pendulum.parse(
"2030-01-01T00:00:00+00:00"
)
assert second.start_trigger_args.trigger_kwargs["moment"] == pendulum.parse(
"2040-06-06T00:00:00+00:00"
)
# the class level template must survive untouched for the next task built from it
assert DateTimeSensorAsync.start_trigger_args.trigger_kwargs == {
"moment": "",
"end_from_trigger": False,
}
25 changes: 25 additions & 0 deletions providers/standard/tests/unit/standard/sensors/test_filesystem.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,3 +238,28 @@ def test_task_defer(self):
task.execute({})

assert isinstance(exc.value.trigger, FileTrigger), "Trigger is not a FileTrigger"

def test_start_trigger_args_are_not_shared_between_tasks(self):
"""Each task must carry its own trigger arguments.

``start_trigger_args`` is a class attribute, so assigning through it made every task
built from this operator advertise the path and timeout of whichever was constructed
last.
"""
with DAG(
dag_id="test_start_trigger_args_not_shared",
schedule=None,
start_date=datetime(2020, 1, 1),
):
first = FileSensor(
task_id="first", filepath="first.txt", deferrable=True, start_from_trigger=True, timeout=60
)
second = FileSensor(
task_id="second", filepath="second.txt", deferrable=True, start_from_trigger=True, timeout=999
)

assert first.start_trigger_args is not second.start_trigger_args
assert first.start_trigger_args.trigger_kwargs["filepath"] == first.path
assert second.start_trigger_args.trigger_kwargs["filepath"] == second.path
assert first.start_trigger_args.timeout == timedelta(seconds=60)
assert second.start_trigger_args.timeout == timedelta(seconds=999)
19 changes: 19 additions & 0 deletions providers/standard/tests/unit/standard/sensors/test_time.py
Original file line number Diff line number Diff line change
Expand Up @@ -135,3 +135,22 @@ def test_execute_complete_accepts_event(self):
op.execute_complete(context={}, event={"status": "success"})
except TypeError as e:
pytest.fail(f"TypeError raised: {e}")

def test_start_trigger_args_are_not_shared_between_tasks(self):
"""Each task must carry its own trigger arguments.

``start_trigger_args`` is a class attribute, so assigning through it made every task
built from this operator advertise the moment of whichever was constructed last.
"""
with DAG(
dag_id="test_start_trigger_args_not_shared",
schedule=None,
start_date=datetime(2020, 1, 1),
):
early = TimeSensor(task_id="early", target_time=time(1, 0), start_from_trigger=True)
late = TimeSensor(task_id="late", target_time=time(23, 0), start_from_trigger=True)

assert early.start_trigger_args is not late.start_trigger_args
assert early.start_trigger_args.trigger_kwargs["moment"] == early.target_datetime
assert late.start_trigger_args.trigger_kwargs["moment"] == late.target_datetime
assert early.target_datetime != late.target_datetime
Loading