Skip to content

Commit a3741e0

Browse files
authored
Optimize deferrable mode execution (#30920)
1 parent 0c28ed0 commit a3741e0

2 files changed

Lines changed: 34 additions & 18 deletions

File tree

  • airflow/providers/google/cloud/sensors
  • tests/providers/google/cloud/sensors

airflow/providers/google/cloud/sensors/gcs.py

Lines changed: 15 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -231,20 +231,21 @@ def execute(self, context: Context) -> None:
231231
if self.deferrable is False:
232232
super().execute(context)
233233
else:
234-
self.defer(
235-
timeout=timedelta(seconds=self.timeout),
236-
trigger=GCSCheckBlobUpdateTimeTrigger(
237-
bucket=self.bucket,
238-
object_name=self.object,
239-
target_date=self.ts_func(context),
240-
poke_interval=self.poke_interval,
241-
google_cloud_conn_id=self.google_cloud_conn_id,
242-
hook_params={
243-
"impersonation_chain": self.impersonation_chain,
244-
},
245-
),
246-
method_name="execute_complete",
247-
)
234+
if not self.poke(context=context):
235+
self.defer(
236+
timeout=timedelta(seconds=self.timeout),
237+
trigger=GCSCheckBlobUpdateTimeTrigger(
238+
bucket=self.bucket,
239+
object_name=self.object,
240+
target_date=self.ts_func(context),
241+
poke_interval=self.poke_interval,
242+
google_cloud_conn_id=self.google_cloud_conn_id,
243+
hook_params={
244+
"impersonation_chain": self.impersonation_chain,
245+
},
246+
),
247+
method_name="execute_complete",
248+
)
248249

249250
def execute_complete(self, context: dict[str, Any], event: dict[str, str] | None = None) -> str:
250251
"""Callback for when the trigger fires."""

tests/providers/google/cloud/sensors/test_gcs.py

Lines changed: 19 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,6 @@
3939
GCSCheckBlobUpdateTimeTrigger,
4040
GCSPrefixBlobTrigger,
4141
)
42-
from tests.providers.google.cloud.utils.airflow_util import create_context
4342

4443
TEST_BUCKET = "TEST_BUCKET"
4544

@@ -247,6 +246,21 @@ def test_should_pass_argument_to_hook(self, mock_hook):
247246
mock_hook.return_value.is_updated_after.assert_called_once_with(TEST_BUCKET, TEST_OBJECT, mock.ANY)
248247
assert result is True
249248

249+
@mock.patch("airflow.providers.google.cloud.sensors.gcs.GCSHook")
250+
@mock.patch("airflow.providers.google.cloud.sensors.gcs.GCSObjectUpdateSensor.defer")
251+
def test_gcs_object_update_sensor_finish_before_deferred(self, mock_defer, mock_hook):
252+
task = GCSObjectUpdateSensor(
253+
task_id="task-id",
254+
bucket=TEST_BUCKET,
255+
object=TEST_OBJECT,
256+
google_cloud_conn_id=TEST_GCP_CONN_ID,
257+
impersonation_chain=TEST_IMPERSONATION_CHAIN,
258+
deferrable=True,
259+
)
260+
mock_hook.return_value.is_updated_after.return_value = True
261+
task.execute(mock.MagicMock())
262+
assert not mock_defer.called
263+
250264

251265
class TestGCSObjectUpdateSensorAsync:
252266
OPERATOR = GCSObjectUpdateSensor(
@@ -257,14 +271,15 @@ class TestGCSObjectUpdateSensorAsync:
257271
deferrable=True,
258272
)
259273

260-
def test_gcs_object_update_sensor_async(self, context):
274+
@mock.patch("airflow.providers.google.cloud.sensors.gcs.GCSHook")
275+
def test_gcs_object_update_sensor_async(self, mock_hook):
261276
"""
262277
Asserts that a task is deferred and a GCSBlobTrigger will be fired
263278
when the GCSObjectUpdateSensorAsync is executed.
264279
"""
265-
280+
mock_hook.return_value.is_updated_after.return_value = False
266281
with pytest.raises(TaskDeferred) as exc:
267-
self.OPERATOR.execute(create_context(self.OPERATOR))
282+
self.OPERATOR.execute(mock.MagicMock())
268283
assert isinstance(
269284
exc.value.trigger, GCSCheckBlobUpdateTimeTrigger
270285
), "Trigger is not a GCSCheckBlobUpdateTimeTrigger"

0 commit comments

Comments
 (0)