From edca154bcd19b2865adc4b85c10c294b052ba895 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sat, 7 Sep 2024 07:27:46 +0100 Subject: [PATCH 1/5] fix consistent return response pubsubsensor --- .../providers/google/cloud/triggers/pubsub.py | 9 ++-- .../google/cloud/triggers/test_pubsub.py | 54 +++++++++++++++++++ 2 files changed, 60 insertions(+), 3 deletions(-) diff --git a/airflow/providers/google/cloud/triggers/pubsub.py b/airflow/providers/google/cloud/triggers/pubsub.py index 535bfe2ba1c68..a4eff23de4a63 100644 --- a/airflow/providers/google/cloud/triggers/pubsub.py +++ b/airflow/providers/google/cloud/triggers/pubsub.py @@ -21,12 +21,12 @@ import asyncio from typing import TYPE_CHECKING, Any, AsyncIterator, Callable, Sequence +from google.cloud.pubsub_v1.types import ReceivedMessage + from airflow.providers.google.cloud.hooks.pubsub import PubSubAsyncHook from airflow.triggers.base import BaseTrigger, TriggerEvent if TYPE_CHECKING: - from google.cloud.pubsub_v1.types import ReceivedMessage - from airflow.utils.context import Context @@ -106,7 +106,10 @@ async def run(self) -> AsyncIterator[TriggerEvent]: # type: ignore[override] ): if self.ack_messages: await self.message_acknowledgement(pulled_messages) - yield TriggerEvent({"status": "success", "message": pulled_messages}) + + messages_json = [ReceivedMessage.to_dict(m) for m in pulled_messages] + + yield TriggerEvent({"status": "success", "message": messages_json}) return self.log.info("Sleeping for %s seconds.", self.poke_interval) await asyncio.sleep(self.poke_interval) diff --git a/tests/providers/google/cloud/triggers/test_pubsub.py b/tests/providers/google/cloud/triggers/test_pubsub.py index d2294eb61414b..180618864bbe8 100644 --- a/tests/providers/google/cloud/triggers/test_pubsub.py +++ b/tests/providers/google/cloud/triggers/test_pubsub.py @@ -16,9 +16,13 @@ # under the License. from __future__ import annotations +from unittest import mock + import pytest +from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.providers.google.cloud.triggers.pubsub import PubsubPullTrigger +from airflow.triggers.base import TriggerEvent TEST_POLL_INTERVAL = 10 TEST_GCP_CONN_ID = "google_cloud_default" @@ -41,6 +45,19 @@ def trigger(): ) +async def generate_messages(count): + return [ + ReceivedMessage( + ack_id=f"{i}", + message={ + "data": f"Message {i}".encode(), + "attributes": {"type": "generated message"}, + }, + ) + for i in range(1, count + 1) + ] + + class TestPubsubPullTrigger: def test_async_pubsub_pull_trigger_serialization_should_execute_successfully(self, trigger): """ @@ -59,3 +76,40 @@ def test_async_pubsub_pull_trigger_serialization_should_execute_successfully(sel "gcp_conn_id": TEST_GCP_CONN_ID, "impersonation_chain": None, } + + @pytest.mark.asyncio + @mock.patch("airflow.providers.google.cloud.hooks.pubsub.PubSubAsyncHook.pull") + async def test_async_pubsub_pull_trigger_return_event(self, mock_pull): + mock_pull.return_value = generate_messages(1) + trigger = PubsubPullTrigger( + project_id=PROJECT_ID, + subscription="subscription", + max_messages=MAX_MESSAGES, + ack_messages=False, + messages_callback=None, + poke_interval=TEST_POLL_INTERVAL, + gcp_conn_id=TEST_GCP_CONN_ID, + impersonation_chain=None, + ) + + expected_event = TriggerEvent( + { + "status": "success", + "message": [ + { + "ack_id": "1", + "message": { + "data": "TWVzc2FnZSAx", + "attributes": {"type": "generated message"}, + "message_id": "", + "ordering_key": "", + }, + "delivery_attempt": 0, + } + ], + } + ) + + response = await trigger.run().asend(None) + + assert response == expected_event From 49e83abe72824a5e9e872f1c30f510524d7c8784 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sat, 14 Sep 2024 12:42:37 +0100 Subject: [PATCH 2/5] removed messages_callback argument to pubsub trigger and using it in execute_complete --- .../providers/google/cloud/sensors/pubsub.py | 24 ++++++++- .../providers/google/cloud/triggers/pubsub.py | 13 +---- .../google/cloud/sensors/test_pubsub.py | 51 +++++++++++++++++++ .../google/cloud/triggers/test_pubsub.py | 3 -- 4 files changed, 74 insertions(+), 17 deletions(-) diff --git a/airflow/providers/google/cloud/sensors/pubsub.py b/airflow/providers/google/cloud/sensors/pubsub.py index cb224d42979b7..e3eb3a3ab85bd 100644 --- a/airflow/providers/google/cloud/sensors/pubsub.py +++ b/airflow/providers/google/cloud/sensors/pubsub.py @@ -22,6 +22,7 @@ from datetime import timedelta from typing import TYPE_CHECKING, Any, Callable, Sequence +from google.cloud import pubsub_v1 from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.configuration import conf @@ -170,7 +171,6 @@ def execute(self, context: Context) -> None: subscription=self.subscription, max_messages=self.max_messages, ack_messages=self.ack_messages, - messages_callback=self.messages_callback, poke_interval=self.poke_interval, gcp_conn_id=self.gcp_conn_id, impersonation_chain=self.impersonation_chain, @@ -178,14 +178,34 @@ def execute(self, context: Context) -> None: method_name="execute_complete", ) - def execute_complete(self, context: dict[str, Any], event: dict[str, str | list[str]]) -> str | list[str]: + def execute_complete(self, context: Context, event: dict[str, str | list[str]]) -> Any: """Return immediately and relies on trigger to throw a success event. Callback for the trigger.""" if event["status"] == "success": self.log.info("Sensor pulls messages: %s", event["message"]) + if self.messages_callback: + message = self._convert_to_received_message(event["message"]) + message_callback_response = self.messages_callback(message, context) + + return message_callback_response + return event["message"] self.log.info("Sensor failed: %s", event["message"]) raise AirflowException(event["message"]) + def _convert_to_received_message(self, messages: Any): + try: + received_messages = [] + for msg in messages: + received_message = pubsub_v1.types.ReceivedMessage(msg) + + received_messages.append(received_message) + + return received_messages + except Exception as e: + raise AirflowException( + f"Error converting triggerer event message back to received message format: {e}" + ) + def _default_message_callback( self, pulled_messages: list[ReceivedMessage], diff --git a/airflow/providers/google/cloud/triggers/pubsub.py b/airflow/providers/google/cloud/triggers/pubsub.py index a4eff23de4a63..db3fe409e942b 100644 --- a/airflow/providers/google/cloud/triggers/pubsub.py +++ b/airflow/providers/google/cloud/triggers/pubsub.py @@ -19,16 +19,13 @@ from __future__ import annotations import asyncio -from typing import TYPE_CHECKING, Any, AsyncIterator, Callable, Sequence +from typing import Any, AsyncIterator, Sequence from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.providers.google.cloud.hooks.pubsub import PubSubAsyncHook from airflow.triggers.base import BaseTrigger, TriggerEvent -if TYPE_CHECKING: - from airflow.utils.context import Context - class PubsubPullTrigger(BaseTrigger): """ @@ -41,11 +38,6 @@ class PubsubPullTrigger(BaseTrigger): :param ack_messages: If True, each message will be acknowledged immediately rather than by any downstream tasks :param gcp_conn_id: Reference to google cloud connection id - :param messages_callback: (Optional) Callback to process received messages. - Its return value will be saved to XCom. - If you are pulling large messages, you probably want to provide a custom callback. - If not provided, the default implementation will convert `ReceivedMessage` objects - into JSON-serializable dicts using `google.protobuf.json_format.MessageToDict` function. :param poke_interval: polling period in seconds to check for the status :param impersonation_chain: Optional service account to impersonate using short-term credentials, or chained list of accounts required to get the access_token @@ -64,7 +56,6 @@ def __init__( max_messages: int, ack_messages: bool, gcp_conn_id: str, - messages_callback: Callable[[list[ReceivedMessage], Context], Any] | None = None, poke_interval: float = 10.0, impersonation_chain: str | Sequence[str] | None = None, ): @@ -73,7 +64,6 @@ def __init__( self.subscription = subscription self.max_messages = max_messages self.ack_messages = ack_messages - self.messages_callback = messages_callback self.poke_interval = poke_interval self.gcp_conn_id = gcp_conn_id self.impersonation_chain = impersonation_chain @@ -88,7 +78,6 @@ def serialize(self) -> tuple[str, dict[str, Any]]: "subscription": self.subscription, "max_messages": self.max_messages, "ack_messages": self.ack_messages, - "messages_callback": self.messages_callback, "poke_interval": self.poke_interval, "gcp_conn_id": self.gcp_conn_id, "impersonation_chain": self.impersonation_chain, diff --git a/tests/providers/google/cloud/sensors/test_pubsub.py b/tests/providers/google/cloud/sensors/test_pubsub.py index a77167dda3037..92b7b6f9c038b 100644 --- a/tests/providers/google/cloud/sensors/test_pubsub.py +++ b/tests/providers/google/cloud/sensors/test_pubsub.py @@ -21,6 +21,7 @@ from unittest import mock import pytest +from google.cloud import pubsub_v1 from google.cloud.pubsub_v1.types import ReceivedMessage from airflow.exceptions import AirflowException, TaskDeferred @@ -197,3 +198,53 @@ def test_pubsub_pull_sensor_async_execute_complete(self): with mock.patch.object(operator.log, "info") as mock_log_info: operator.execute_complete(context={}, event={"status": "success", "message": test_message}) mock_log_info.assert_called_with("Sensor pulls messages: %s", test_message) + + @mock.patch("airflow.providers.google.cloud.sensors.pubsub.PubSubHook") + def test_pubsub_pull_sensor_async_execute_complete_use_message_callback(self, mock_hook): + test_message = [ + { + "ack_id": "UAYWLF1GSFE3GQhoUQ5PXiM_NSAoRRIJB08CKF15MU0sQVhwaFENGXJ9YHxrUxsDV0ECel1RGQdoTm11H4GglfRLQ1RrWBIHB01Vel5TEwxoX11wBnm4vPO6v8vgfwk9OpX-8tltO6ywsP9GZiM9XhJLLD5-LzlFQV5AEkwkDERJUytDCypYEU4EISE-MD5FU0Q", + "message": { + "data": "aGkgZnJvbSBjbG91ZCBjb25zb2xlIQ==", + "message_id": "12165864188103151", + "publish_time": "2024-08-28T11:49:50.962Z", + "attributes": {}, + "ordering_key": "", + }, + "delivery_attempt": 0, + } + ] + received_message_format = [] + for msg in test_message: + received_message_format.append(pubsub_v1.types.ReceivedMessage(msg)) + + messages_callback_return_value = "custom_message_from_callback" + + def messages_callback( + pulled_messages: list[ReceivedMessage], + context: dict[str, Any], + ): + assert pulled_messages == received_message_format + + assert isinstance(context, dict) + for key in context.keys(): + assert isinstance(key, str) + + return messages_callback_return_value + + messages_callback = mock.Mock(side_effect=messages_callback) + + operator = PubSubPullSensor( + task_id="test_task", + ack_messages=True, + project_id=TEST_PROJECT, + subscription=TEST_SUBSCRIPTION, + deferrable=True, + messages_callback=messages_callback, + ) + mock_hook.return_value.pull.return_value = received_message_format + + with mock.patch.object(operator.log, "info") as mock_log_info: + resp = operator.execute_complete(context={}, event={"status": "success", "message": test_message}) + mock_log_info.assert_called_with("Sensor pulls messages: %s", test_message) + assert resp == messages_callback_return_value diff --git a/tests/providers/google/cloud/triggers/test_pubsub.py b/tests/providers/google/cloud/triggers/test_pubsub.py index 180618864bbe8..f8779ac1ef10d 100644 --- a/tests/providers/google/cloud/triggers/test_pubsub.py +++ b/tests/providers/google/cloud/triggers/test_pubsub.py @@ -38,7 +38,6 @@ def trigger(): subscription="subscription", max_messages=MAX_MESSAGES, ack_messages=ACK_MESSAGES, - messages_callback=None, poke_interval=TEST_POLL_INTERVAL, gcp_conn_id=TEST_GCP_CONN_ID, impersonation_chain=None, @@ -71,7 +70,6 @@ def test_async_pubsub_pull_trigger_serialization_should_execute_successfully(sel "subscription": "subscription", "max_messages": MAX_MESSAGES, "ack_messages": ACK_MESSAGES, - "messages_callback": None, "poke_interval": TEST_POLL_INTERVAL, "gcp_conn_id": TEST_GCP_CONN_ID, "impersonation_chain": None, @@ -86,7 +84,6 @@ async def test_async_pubsub_pull_trigger_return_event(self, mock_pull): subscription="subscription", max_messages=MAX_MESSAGES, ack_messages=False, - messages_callback=None, poke_interval=TEST_POLL_INTERVAL, gcp_conn_id=TEST_GCP_CONN_ID, impersonation_chain=None, From 111c6327bb5c2b4dd4421c8557b8b7de670a01e3 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Sat, 14 Sep 2024 13:00:39 +0100 Subject: [PATCH 3/5] updated variable name --- airflow/providers/google/cloud/sensors/pubsub.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/airflow/providers/google/cloud/sensors/pubsub.py b/airflow/providers/google/cloud/sensors/pubsub.py index e3eb3a3ab85bd..69225db7e3491 100644 --- a/airflow/providers/google/cloud/sensors/pubsub.py +++ b/airflow/providers/google/cloud/sensors/pubsub.py @@ -183,8 +183,8 @@ def execute_complete(self, context: Context, event: dict[str, str | list[str]]) if event["status"] == "success": self.log.info("Sensor pulls messages: %s", event["message"]) if self.messages_callback: - message = self._convert_to_received_message(event["message"]) - message_callback_response = self.messages_callback(message, context) + received_message_fmt = self._convert_to_received_message(event["message"]) + message_callback_response = self.messages_callback(received_message_fmt, context) return message_callback_response From c2fee3ddb50174a4d6268b95d36b48af09ddd6b2 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Mon, 23 Sep 2024 04:38:05 +0100 Subject: [PATCH 4/5] updates as per comments, added return types and refactored logic --- airflow/providers/google/cloud/sensors/pubsub.py | 11 +++-------- 1 file changed, 3 insertions(+), 8 deletions(-) diff --git a/airflow/providers/google/cloud/sensors/pubsub.py b/airflow/providers/google/cloud/sensors/pubsub.py index 69225db7e3491..d240247684927 100644 --- a/airflow/providers/google/cloud/sensors/pubsub.py +++ b/airflow/providers/google/cloud/sensors/pubsub.py @@ -179,7 +179,7 @@ def execute(self, context: Context) -> None: ) def execute_complete(self, context: Context, event: dict[str, str | list[str]]) -> Any: - """Return immediately and relies on trigger to throw a success event. Callback for the trigger.""" + """If messages_callback is provided, execute it; otherwise, return immediately with trigger event message.""" if event["status"] == "success": self.log.info("Sensor pulls messages: %s", event["message"]) if self.messages_callback: @@ -192,14 +192,9 @@ def execute_complete(self, context: Context, event: dict[str, str | list[str]]) self.log.info("Sensor failed: %s", event["message"]) raise AirflowException(event["message"]) - def _convert_to_received_message(self, messages: Any): + def _convert_to_received_message(self, messages: Any) -> list[Any]: try: - received_messages = [] - for msg in messages: - received_message = pubsub_v1.types.ReceivedMessage(msg) - - received_messages.append(received_message) - + received_messages = [pubsub_v1.types.ReceivedMessage(msg) for msg in messages] return received_messages except Exception as e: raise AirflowException( From 0f718a411f9e98de98247694dcde8e0e97007073 Mon Sep 17 00:00:00 2001 From: Pavan Kumar Date: Mon, 23 Sep 2024 10:26:06 +0100 Subject: [PATCH 5/5] update types, tests and use inherit exception --- airflow/providers/google/cloud/sensors/pubsub.py | 15 +++++++++------ .../providers/google/cloud/sensors/test_pubsub.py | 11 ++++------- .../google/cloud/triggers/test_pubsub.py | 2 +- 3 files changed, 14 insertions(+), 14 deletions(-) diff --git a/airflow/providers/google/cloud/sensors/pubsub.py b/airflow/providers/google/cloud/sensors/pubsub.py index d240247684927..aa74411f072e5 100644 --- a/airflow/providers/google/cloud/sensors/pubsub.py +++ b/airflow/providers/google/cloud/sensors/pubsub.py @@ -35,6 +35,10 @@ from airflow.utils.context import Context +class PubSubMessageTransformException(AirflowException): + """Raise when messages failed to convert pubsub received format.""" + + class PubSubPullSensor(BaseSensorOperator): """ Pulls messages from a PubSub subscription and passes them through XCom. @@ -183,21 +187,20 @@ def execute_complete(self, context: Context, event: dict[str, str | list[str]]) if event["status"] == "success": self.log.info("Sensor pulls messages: %s", event["message"]) if self.messages_callback: - received_message_fmt = self._convert_to_received_message(event["message"]) - message_callback_response = self.messages_callback(received_message_fmt, context) - - return message_callback_response + received_messages = self._convert_to_received_messages(event["message"]) + _return_value = self.messages_callback(received_messages, context) + return _return_value return event["message"] self.log.info("Sensor failed: %s", event["message"]) raise AirflowException(event["message"]) - def _convert_to_received_message(self, messages: Any) -> list[Any]: + def _convert_to_received_messages(self, messages: Any) -> list[ReceivedMessage]: try: received_messages = [pubsub_v1.types.ReceivedMessage(msg) for msg in messages] return received_messages except Exception as e: - raise AirflowException( + raise PubSubMessageTransformException( f"Error converting triggerer event message back to received message format: {e}" ) diff --git a/tests/providers/google/cloud/sensors/test_pubsub.py b/tests/providers/google/cloud/sensors/test_pubsub.py index 92b7b6f9c038b..5a3fb170b7482 100644 --- a/tests/providers/google/cloud/sensors/test_pubsub.py +++ b/tests/providers/google/cloud/sensors/test_pubsub.py @@ -214,9 +214,8 @@ def test_pubsub_pull_sensor_async_execute_complete_use_message_callback(self, mo "delivery_attempt": 0, } ] - received_message_format = [] - for msg in test_message: - received_message_format.append(pubsub_v1.types.ReceivedMessage(msg)) + + received_messages = [pubsub_v1.types.ReceivedMessage(msg) for msg in test_message] messages_callback_return_value = "custom_message_from_callback" @@ -224,7 +223,7 @@ def messages_callback( pulled_messages: list[ReceivedMessage], context: dict[str, Any], ): - assert pulled_messages == received_message_format + assert pulled_messages == received_messages assert isinstance(context, dict) for key in context.keys(): @@ -232,8 +231,6 @@ def messages_callback( return messages_callback_return_value - messages_callback = mock.Mock(side_effect=messages_callback) - operator = PubSubPullSensor( task_id="test_task", ack_messages=True, @@ -242,7 +239,7 @@ def messages_callback( deferrable=True, messages_callback=messages_callback, ) - mock_hook.return_value.pull.return_value = received_message_format + mock_hook.return_value.pull.return_value = received_messages with mock.patch.object(operator.log, "info") as mock_log_info: resp = operator.execute_complete(context={}, event={"status": "success", "message": test_message}) diff --git a/tests/providers/google/cloud/triggers/test_pubsub.py b/tests/providers/google/cloud/triggers/test_pubsub.py index f8779ac1ef10d..e1a4e178d2918 100644 --- a/tests/providers/google/cloud/triggers/test_pubsub.py +++ b/tests/providers/google/cloud/triggers/test_pubsub.py @@ -44,7 +44,7 @@ def trigger(): ) -async def generate_messages(count): +async def generate_messages(count: int) -> list[ReceivedMessage]: return [ ReceivedMessage( ack_id=f"{i}",