This is an automated email from the ASF dual-hosted git repository.
shahar1 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 11a3892daa9 Fix Dataflow sensors losing their XCom value when not
deferred (#71086)
11a3892daa9 is described below
commit 11a3892daa9dfd5e84fb037d49373b174a372fb8
Author: PoAn Yang <[email protected]>
AuthorDate: Thu Sep 17 20:16:49 2026 +0900
Fix Dataflow sensors losing their XCom value when not deferred (#71086)
---
.../providers/google/cloud/sensors/dataflow.py | 99 ++++++++++++----------
.../example_dataflow_native_python_async.py | 19 +++++
.../unit/google/cloud/sensors/test_dataflow.py | 53 +++++++++++-
3 files changed, 124 insertions(+), 47 deletions(-)
diff --git
a/providers/google/src/airflow/providers/google/cloud/sensors/dataflow.py
b/providers/google/src/airflow/providers/google/cloud/sensors/dataflow.py
index 1b4f74bb01b..4d149db7af5 100644
--- a/providers/google/src/airflow/providers/google/cloud/sensors/dataflow.py
+++ b/providers/google/src/airflow/providers/google/cloud/sensors/dataflow.py
@@ -219,7 +219,7 @@ class DataflowJobMetricsSensor(BaseSensorOperator):
self.deferrable = deferrable
self.poll_interval = poll_interval
- def poke(self, context: Context) -> bool:
+ def poke(self, context: Context) -> PokeReturnValue | bool:
if self.fail_on_terminal_state:
job = self.hook.get_job(
job_id=self.job_id,
@@ -236,26 +236,35 @@ class DataflowJobMetricsSensor(BaseSensorOperator):
project_id=self.project_id,
location=self.location,
)
- return result["metrics"] if self.callback is None else
self.callback(result["metrics"])
+ result = result["metrics"] if self.callback is None else
self.callback(result["metrics"])
+
+ if isinstance(result, PokeReturnValue):
+ return result
+
+ if bool(result):
+ return PokeReturnValue(
+ is_done=True,
+ xcom_value=result,
+ )
+ return False
def execute(self, context: Context) -> Any:
"""Airflow runs this method on the worker and defers using the
trigger."""
if not self.deferrable:
- super().execute(context)
- else:
- self.defer(
- timeout=self.execution_timeout,
- trigger=DataflowJobMetricsTrigger(
- job_id=self.job_id,
- project_id=self.project_id,
- location=self.location,
- gcp_conn_id=self.gcp_conn_id,
- poll_sleep=self.poll_interval,
- impersonation_chain=self.impersonation_chain,
- fail_on_terminal_state=self.fail_on_terminal_state,
- ),
- method_name="execute_complete",
- )
+ return super().execute(context)
+ self.defer(
+ timeout=self.execution_timeout,
+ trigger=DataflowJobMetricsTrigger(
+ job_id=self.job_id,
+ project_id=self.project_id,
+ location=self.location,
+ gcp_conn_id=self.gcp_conn_id,
+ poll_sleep=self.poll_interval,
+ impersonation_chain=self.impersonation_chain,
+ fail_on_terminal_state=self.fail_on_terminal_state,
+ ),
+ method_name="execute_complete",
+ )
def execute_complete(self, context: Context, event: dict[str, str | list])
-> Any:
"""
@@ -372,21 +381,20 @@ class DataflowJobMessagesSensor(BaseSensorOperator):
def execute(self, context: Context) -> Any:
"""Airflow runs this method on the worker and defers using the
trigger."""
if not self.deferrable:
- super().execute(context)
- else:
- self.defer(
- timeout=self.execution_timeout,
- trigger=DataflowJobMessagesTrigger(
- job_id=self.job_id,
- project_id=self.project_id,
- location=self.location,
- gcp_conn_id=self.gcp_conn_id,
- poll_sleep=self.poll_interval,
- impersonation_chain=self.impersonation_chain,
- fail_on_terminal_state=self.fail_on_terminal_state,
- ),
- method_name="execute_complete",
- )
+ return super().execute(context)
+ self.defer(
+ timeout=self.execution_timeout,
+ trigger=DataflowJobMessagesTrigger(
+ job_id=self.job_id,
+ project_id=self.project_id,
+ location=self.location,
+ gcp_conn_id=self.gcp_conn_id,
+ poll_sleep=self.poll_interval,
+ impersonation_chain=self.impersonation_chain,
+ fail_on_terminal_state=self.fail_on_terminal_state,
+ ),
+ method_name="execute_complete",
+ )
def execute_complete(self, context: Context, event: dict[str, str | list])
-> Any:
"""
@@ -502,20 +510,19 @@ class
DataflowJobAutoScalingEventsSensor(BaseSensorOperator):
def execute(self, context: Context) -> Any:
"""Airflow runs this method on the worker and defers using the
trigger."""
if not self.deferrable:
- super().execute(context)
- else:
- self.defer(
- trigger=DataflowJobAutoScalingEventTrigger(
- job_id=self.job_id,
- project_id=self.project_id,
- location=self.location,
- gcp_conn_id=self.gcp_conn_id,
- poll_sleep=self.poll_interval,
- impersonation_chain=self.impersonation_chain,
- fail_on_terminal_state=self.fail_on_terminal_state,
- ),
- method_name="execute_complete",
- )
+ return super().execute(context)
+ self.defer(
+ trigger=DataflowJobAutoScalingEventTrigger(
+ job_id=self.job_id,
+ project_id=self.project_id,
+ location=self.location,
+ gcp_conn_id=self.gcp_conn_id,
+ poll_sleep=self.poll_interval,
+ impersonation_chain=self.impersonation_chain,
+ fail_on_terminal_state=self.fail_on_terminal_state,
+ ),
+ method_name="execute_complete",
+ )
def execute_complete(self, context: Context, event: dict[str, str | list])
-> Any:
"""
diff --git
a/providers/google/tests/system/google/cloud/dataflow/example_dataflow_native_python_async.py
b/providers/google/tests/system/google/cloud/dataflow/example_dataflow_native_python_async.py
index d5d24283859..fd60ed08ed7 100644
---
a/providers/google/tests/system/google/cloud/dataflow/example_dataflow_native_python_async.py
+++
b/providers/google/tests/system/google/cloud/dataflow/example_dataflow_native_python_async.py
@@ -39,6 +39,7 @@ from airflow.providers.google.cloud.sensors.dataflow import (
DataflowJobMetricsSensor,
DataflowJobStatusSensor,
)
+from airflow.providers.standard.operators.python import PythonOperator
try:
from airflow.sdk import TriggerRule
@@ -66,6 +67,17 @@ default_args = {
}
log = logging.getLogger(__name__)
+
+def _assert_sensors_pushed_xcom(ti):
+ """Check that each sensor pushed its callback result to XCom while running
in poke mode."""
+ for task_id in (
+ "wait_for_python_job_async_metric",
+ "wait_for_python_job_async_message",
+ "wait_for_python_job_async_autoscaling_event",
+ ):
+ assert ti.xcom_pull(task_ids=task_id) is not None, f"{task_id} did not
push a value to XCom"
+
+
with DAG(
DAG_ID,
default_args=default_args,
@@ -84,6 +96,8 @@ with DAG(
py_options=[],
pipeline_options={
"output": GCS_OUTPUT,
+ "machine_type": "e2-standard-2",
+ "worker_zone": "europe-west3-a",
},
py_requirements=["apache-beam[gcp]==2.67.0"],
py_interpreter="python3",
@@ -165,6 +179,10 @@ with DAG(
)
# [END howto_sensor_wait_for_job_autoscaling_event]
+ assert_sensors_pushed_xcom = PythonOperator(
+ task_id="assert_sensors_pushed_xcom",
python_callable=_assert_sensors_pushed_xcom
+ )
+
delete_bucket = GCSDeleteBucketOperator(
task_id="delete_bucket", bucket_name=BUCKET_NAME,
trigger_rule=TriggerRule.ALL_DONE
)
@@ -180,6 +198,7 @@ with DAG(
wait_for_python_job_async_message,
wait_for_python_job_async_autoscaling_event,
]
+ >> assert_sensors_pushed_xcom
# TEST TEARDOWN
>> delete_bucket
)
diff --git a/providers/google/tests/unit/google/cloud/sensors/test_dataflow.py
b/providers/google/tests/unit/google/cloud/sensors/test_dataflow.py
index 21cdca18328..b77a671c47e 100644
--- a/providers/google/tests/unit/google/cloud/sensors/test_dataflow.py
+++ b/providers/google/tests/unit/google/cloud/sensors/test_dataflow.py
@@ -197,7 +197,7 @@ class TestDataflowJobMetricsSensor:
mock_get_job.return_value = {"id": TEST_JOB_ID, "currentState":
job_current_state}
results = task.poke(mock.MagicMock())
- assert callback.return_value == results
+ assert callback.return_value == results.xcom_value
mock_hook.assert_called_once_with(
gcp_conn_id=TEST_GCP_CONN_ID,
@@ -241,6 +241,23 @@ class TestDataflowJobMetricsSensor:
mock_fetch_job_messages_by_id.assert_not_called()
callback.assert_not_called()
+ @mock.patch("airflow.providers.google.cloud.sensors.dataflow.DataflowHook")
+ def test_execute_returns_xcom_value_in_non_deferrable_mode(self,
mock_hook):
+ """Deferrable mode returns the metrics through execute_complete; poke
mode must match it."""
+ callback = mock.MagicMock()
+ task = DataflowJobMetricsSensor(
+ task_id=TEST_TASK_ID,
+ job_id=TEST_JOB_ID,
+ callback=callback,
+ fail_on_terminal_state=False,
+ location=TEST_LOCATION,
+ project_id=TEST_PROJECT_ID,
+ gcp_conn_id=TEST_GCP_CONN_ID,
+ impersonation_chain=TEST_IMPERSONATION_CHAIN,
+ )
+
+ assert task.execute(mock.MagicMock()) == callback.return_value
+
@mock.patch("airflow.providers.google.cloud.hooks.dataflow.AsyncDataflowHook")
def test_execute_enters_deferred_state(self, mock_hook):
"""
@@ -419,6 +436,23 @@ class TestDataflowJobMessagesSensor:
mock_fetch_job_messages_by_id.assert_not_called()
callback.assert_not_called()
+ @mock.patch("airflow.providers.google.cloud.sensors.dataflow.DataflowHook")
+ def test_execute_returns_xcom_value_in_non_deferrable_mode(self,
mock_hook):
+ """Deferrable mode returns the messages through execute_complete; poke
mode must match it."""
+ callback = mock.MagicMock()
+ task = DataflowJobMessagesSensor(
+ task_id=TEST_TASK_ID,
+ job_id=TEST_JOB_ID,
+ callback=callback,
+ fail_on_terminal_state=False,
+ location=TEST_LOCATION,
+ project_id=TEST_PROJECT_ID,
+ gcp_conn_id=TEST_GCP_CONN_ID,
+ impersonation_chain=TEST_IMPERSONATION_CHAIN,
+ )
+
+ assert task.execute(mock.MagicMock()) == callback.return_value
+
@mock.patch("airflow.providers.google.cloud.hooks.dataflow.AsyncDataflowHook")
def test_execute_enters_deferred_state(self, mock_hook):
"""
@@ -595,6 +629,23 @@ class TestDataflowJobAutoScalingEventsSensor:
mock_fetch_job_autoscaling_events_by_id.assert_not_called()
callback.assert_not_called()
+ @mock.patch("airflow.providers.google.cloud.sensors.dataflow.DataflowHook")
+ def test_execute_returns_xcom_value_in_non_deferrable_mode(self,
mock_hook):
+ """Deferrable mode returns the events through execute_complete; poke
mode must match it."""
+ callback = mock.MagicMock()
+ task = DataflowJobAutoScalingEventsSensor(
+ task_id=TEST_TASK_ID,
+ job_id=TEST_JOB_ID,
+ callback=callback,
+ fail_on_terminal_state=False,
+ location=TEST_LOCATION,
+ project_id=TEST_PROJECT_ID,
+ gcp_conn_id=TEST_GCP_CONN_ID,
+ impersonation_chain=TEST_IMPERSONATION_CHAIN,
+ )
+
+ assert task.execute(mock.MagicMock()) == callback.return_value
+
@mock.patch("airflow.providers.google.cloud.hooks.dataflow.AsyncDataflowHook")
def test_execute_enters_deferred_state(self, mock_hook):
"""