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):
         """

Reply via email to