From 58140c701b9dc55a65d36ece7cfd217ed197435d Mon Sep 17 00:00:00 2001 From: Tim Perry Date: Mon, 12 Aug 2024 16:24:33 -0700 Subject: [PATCH 1/4] Added callback to receive message contents --- .../providers/microsoft/azure/hooks/asb.py | 44 ++++++++++-- .../microsoft/azure/operators/asb.py | 26 ++++++- .../microsoft/azure/hooks/test_asb.py | 61 ++++++++++++++++ .../microsoft/azure/operators/test_asb.py | 69 +++++++++++++++++++ 4 files changed, 191 insertions(+), 9 deletions(-) diff --git a/airflow/providers/microsoft/azure/hooks/asb.py b/airflow/providers/microsoft/azure/hooks/asb.py index e70b6d6554c0a..518f244ee42b7 100644 --- a/airflow/providers/microsoft/azure/hooks/asb.py +++ b/airflow/providers/microsoft/azure/hooks/asb.py @@ -16,7 +16,7 @@ # under the License. from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Callable from azure.servicebus import ServiceBusClient, ServiceBusMessage, ServiceBusSender from azure.servicebus.management import QueueProperties, ServiceBusAdministrationClient @@ -28,6 +28,9 @@ get_sync_default_azure_credential, ) +MessageCallback = Callable[[ServiceBusMessage], None] + + if TYPE_CHECKING: from azure.identity import DefaultAzureCredential @@ -270,7 +273,11 @@ def send_batch_message(sender: ServiceBusSender, messages: list[str]): sender.send_messages(batch_message) def receive_message( - self, queue_name, max_message_count: int | None = 1, max_wait_time: float | None = None + self, + queue_name, + max_message_count: int | None = 1, + max_wait_time: float | None = None, + message_callback: MessageCallback | None = None, ): """ Receive a batch of messages at once in a specified Queue name. @@ -278,6 +285,9 @@ def receive_message( :param queue_name: The name of the queue name or a QueueProperties with name. :param max_message_count: Maximum number of messages in the batch. :param max_wait_time: Maximum time to wait in seconds for the first message to arrive. + :param message_callback: Optional callback to process each message. If not provided, then + the message will be logged and completed. If provided, and throws an exception, the + message will be abandoned for future redelivery. """ if queue_name is None: raise TypeError("Queue name cannot be None.") @@ -289,8 +299,7 @@ def receive_message( max_message_count=max_message_count, max_wait_time=max_wait_time ) for msg in received_msgs: - self.log.info(msg) - receiver.complete_message(msg) + self._process_message(msg, message_callback, receiver) def receive_subscription_message( self, @@ -298,6 +307,7 @@ def receive_subscription_message( subscription_name: str, max_message_count: int | None, max_wait_time: float | None, + message_callback: MessageCallback | None = None, ): """ Receive a batch of subscription message at once. @@ -326,5 +336,27 @@ def receive_subscription_message( max_message_count=max_message_count, max_wait_time=max_wait_time ) for msg in received_msgs: - self.log.info(msg) - subscription_receiver.complete_message(msg) + self._process_message(msg, message_callback, subscription_receiver) + + def _process_message(self, msg, message_callback, receiver): + """ + Process the message by calling the message_callback or logging the message. + + :param msg: The message to process. + :param message_callback: Optional callback to process each message. If not provided, then + the message will be logged and completed. If provided, and throws an exception, the + message will be abandoned for future redelivery. + :param receiver: The receiver that received the message. + """ + print("message_callback:", message_callback) + if message_callback is None: + self.log.info(msg) + receiver.complete_message(msg) + else: + try: + message_callback(msg) + receiver.complete_message(msg) + except Exception as e: + self.log.error("Error processing message: %s", e) + receiver.abandon_message(msg) + raise e diff --git a/airflow/providers/microsoft/azure/operators/asb.py b/airflow/providers/microsoft/azure/operators/asb.py index 946f4a7959b9f..85619526cfb93 100644 --- a/airflow/providers/microsoft/azure/operators/asb.py +++ b/airflow/providers/microsoft/azure/operators/asb.py @@ -16,7 +16,7 @@ # under the License. from __future__ import annotations -from typing import TYPE_CHECKING, Any, Sequence +from typing import TYPE_CHECKING, Any, Callable, Sequence from azure.core.exceptions import ResourceNotFoundError @@ -26,10 +26,13 @@ if TYPE_CHECKING: import datetime + from azure.servicebus import ServiceBusMessage from azure.servicebus.management._models import AuthorizationRule from airflow.utils.context import Context + MessageCallback = Callable[[ServiceBusMessage], None] + class AzureServiceBusCreateQueueOperator(BaseOperator): """ @@ -140,6 +143,9 @@ class AzureServiceBusReceiveMessageOperator(BaseOperator): :param max_wait_time: Maximum time to wait in seconds for the first message to arrive. :param azure_service_bus_conn_id: Reference to the :ref: `Azure Service Bus connection `. + :param message_callback: Optional callback to process each message. If not provided, then + the message will be logged and completed. If provided, and throws an exception, the + message will be abandoned for future redelivery. """ template_fields: Sequence[str] = ("queue_name",) @@ -152,6 +158,7 @@ def __init__( azure_service_bus_conn_id: str = "azure_service_bus_default", max_message_count: int = 10, max_wait_time: float = 5, + message_callback: MessageCallback | None = None, **kwargs, ) -> None: super().__init__(**kwargs) @@ -159,6 +166,7 @@ def __init__( self.azure_service_bus_conn_id = azure_service_bus_conn_id self.max_message_count = max_message_count self.max_wait_time = max_wait_time + self.message_callback = message_callback def execute(self, context: Context) -> None: """Receive Message in specific queue in Service Bus namespace by connecting to Service Bus client.""" @@ -167,7 +175,10 @@ def execute(self, context: Context) -> None: # Receive message hook.receive_message( - self.queue_name, max_message_count=self.max_message_count, max_wait_time=self.max_wait_time + self.queue_name, + max_message_count=self.max_message_count, + max_wait_time=self.max_wait_time, + message_callback=self.message_callback, ) @@ -515,6 +526,9 @@ class ASBReceiveSubscriptionMessageOperator(BaseOperator): an empty list will be returned. :param azure_service_bus_conn_id: Reference to the :ref:`Azure Service Bus connection `. + :param message_callback: Optional callback to process each message. If not provided, then + the message will be logged and completed. If provided, and throws an exception, the + message will be abandoned for future redelivery. """ template_fields: Sequence[str] = ("topic_name", "subscription_name") @@ -528,6 +542,7 @@ def __init__( max_message_count: int | None = 1, max_wait_time: float | None = 5, azure_service_bus_conn_id: str = "azure_service_bus_default", + message_callback: MessageCallback | None = None, **kwargs, ) -> None: super().__init__(**kwargs) @@ -536,6 +551,7 @@ def __init__( self.max_message_count = max_message_count self.max_wait_time = max_wait_time self.azure_service_bus_conn_id = azure_service_bus_conn_id + self.message_callback = message_callback def execute(self, context: Context) -> None: """Receive Message in specific queue in Service Bus namespace by connecting to Service Bus client.""" @@ -544,7 +560,11 @@ def execute(self, context: Context) -> None: # Receive message hook.receive_subscription_message( - self.topic_name, self.subscription_name, self.max_message_count, self.max_wait_time + self.topic_name, + self.subscription_name, + self.max_message_count, + self.max_wait_time, + message_callback=self.message_callback, ) diff --git a/tests/providers/microsoft/azure/hooks/test_asb.py b/tests/providers/microsoft/azure/hooks/test_asb.py index eb35447820ae9..5cbdad782cbe0 100644 --- a/tests/providers/microsoft/azure/hooks/test_asb.py +++ b/tests/providers/microsoft/azure/hooks/test_asb.py @@ -16,6 +16,7 @@ # under the License. from __future__ import annotations +from typing import Any from unittest import mock import pytest @@ -265,6 +266,31 @@ def test_receive_message(self, mock_sb_client, mock_service_bus_message): ] mock_sb_client.assert_has_calls(expected_calls) + @mock.patch("azure.servicebus.ServiceBusMessage", autospec=True) + @mock.patch(f"{MODULE}.MessageHook.get_conn", autospec=True) + def test_receive_message_callback(self, mock_sb_client, mock_service_bus_message): + """ + Test `receive_message` hook function and assert the function with mock value, + mock the azure service bus `receive_messages` function + """ + hook = MessageHook(azure_service_bus_conn_id=self.conn_id) + + mock_sb_client.return_value.__enter__.return_value.get_queue_receiver.return_value.__enter__.return_value.receive_messages.return_value = [ + mock_service_bus_message + ] + + received_messages = [] + + def message_callback(msg: Any) -> None: + nonlocal received_messages + print("received message:", msg) + received_messages.append(msg) + + hook.receive_message(self.queue_name, message_callback=message_callback) + + assert len(received_messages) == 1 + assert received_messages[0] == mock_service_bus_message + @mock.patch(f"{MODULE}.MessageHook.get_conn") def test_receive_message_exception(self, mock_sb_client): """ @@ -300,6 +326,41 @@ def test_receive_subscription_message(self, mock_sb_client): ] mock_sb_client.assert_has_calls(expected_calls) + @mock.patch(f"{MODULE}.MessageHook.get_conn") + def test_receive_subscription_message_callback(self, mock_sb_client): + """ + Test `receive_subscription_message` hook function and assert the function with mock value, + mock the azure service bus `receive_message` function of subscription + """ + subscription_name = "subscription_1" + topic_name = "topic_name" + max_message_count = 10 + max_wait_time = 5 + hook = MessageHook(azure_service_bus_conn_id=self.conn_id) + + mock_sb_message0 = ServiceBusMessage("message0") + mock_sb_message1 = ServiceBusMessage("message1") + + mock_sb_client.return_value.__enter__.return_value.get_subscription_receiver.return_value.__enter__.return_value.receive_messages.return_value = [ + mock_sb_message0, + mock_sb_message1, + ] + + received_messages = [] + + def message_callback(msg: ServiceBusMessage) -> None: + nonlocal received_messages + print("received message:", msg) + received_messages.append(msg) + + hook.receive_subscription_message( + topic_name, subscription_name, max_message_count, max_wait_time, message_callback=message_callback + ) + + assert len(received_messages) == 2 + assert received_messages[0] == mock_sb_message0 + assert received_messages[1] == mock_sb_message1 + @pytest.mark.parametrize( "mock_subscription_name, mock_topic_name, mock_max_count, mock_wait_time", [("subscription_1", None, None, None), (None, "topic_1", None, None)], diff --git a/tests/providers/microsoft/azure/operators/test_asb.py b/tests/providers/microsoft/azure/operators/test_asb.py index 774d8a071dd1f..42b770095b4e7 100644 --- a/tests/providers/microsoft/azure/operators/test_asb.py +++ b/tests/providers/microsoft/azure/operators/test_asb.py @@ -210,6 +210,30 @@ def test_receive_message_queue(self, mock_get_conn): ] mock_get_conn.assert_has_calls(expected_calls) + @mock.patch("airflow.providers.microsoft.azure.hooks.asb.MessageHook.get_conn") + def test_receive_message_queue_callback(self, mock_get_conn): + """ + Test AzureServiceBusReceiveMessageOperator by mock connection, values + and the service bus receive message + """ + mock_service_bus_message = ServiceBusMessage("Test message") + mock_get_conn.return_value.__enter__.return_value.get_queue_receiver.return_value.__enter__.return_value.receive_messages.return_value = [ + mock_service_bus_message + ] + + messages_received = [] + + def message_callback(msg): + messages_received.append(msg) + print(msg) + + asb_receive_queue_operator = AzureServiceBusReceiveMessageOperator( + task_id="asb_receive_message_queue", queue_name=QUEUE_NAME, message_callback=message_callback + ) + asb_receive_queue_operator.execute(None) + assert len(messages_received) == 1 + assert messages_received[0] == mock_service_bus_message + class TestABSTopicCreateOperator: def test_init(self): @@ -430,6 +454,51 @@ def test_receive_message_queue(self, mock_get_conn): ] mock_get_conn.assert_has_calls(expected_calls) + @mock.patch("airflow.providers.microsoft.azure.hooks.asb.MessageHook.get_conn") + def test_receive_message_queue_callback(self, mock_get_conn): + """ + Test ASBReceiveSubscriptionMessageOperator by mock connection, values + and the service bus receive message + """ + + mock_sb_message0 = ServiceBusMessage("Test message 0") + mock_sb_message1 = ServiceBusMessage("Test message 1") + mock_get_conn.return_value.__enter__.return_value.get_subscription_receiver.return_value.__enter__.return_value.receive_messages.return_value = [ + mock_sb_message0, + mock_sb_message1, + ] + + messages_received = [] + + def message_callback(msg): + messages_received.append(msg) + print(msg) + + asb_subscription_receive_message = ASBReceiveSubscriptionMessageOperator( + task_id="asb_subscription_receive_message", + topic_name=TOPIC_NAME, + subscription_name=SUBSCRIPTION_NAME, + max_message_count=10, + message_callback=message_callback, + ) + + asb_subscription_receive_message.execute(None) + expected_calls = [ + mock.call() + .__enter__() + .get_subscription_receiver(SUBSCRIPTION_NAME, TOPIC_NAME) + .__enter__() + .receive_messages(max_message_count=10, max_wait_time=5) + .get_subscription_receiver(SUBSCRIPTION_NAME, TOPIC_NAME) + .__exit__() + .mock_call() + .__exit__ + ] + mock_get_conn.assert_has_calls(expected_calls) + assert len(messages_received) == 2 + assert messages_received[0] == mock_sb_message0 + assert messages_received[1] == mock_sb_message1 + class TestASBTopicDeleteOperator: def test_init(self): From f4cc516dfa5f904cac2c676aede96dd03490f61c Mon Sep 17 00:00:00 2001 From: Tim Perry Date: Fri, 30 Aug 2024 08:08:50 -0700 Subject: [PATCH 2/4] remove debug printing --- airflow/providers/microsoft/azure/hooks/asb.py | 1 - 1 file changed, 1 deletion(-) diff --git a/airflow/providers/microsoft/azure/hooks/asb.py b/airflow/providers/microsoft/azure/hooks/asb.py index 518f244ee42b7..00d0994e65dbe 100644 --- a/airflow/providers/microsoft/azure/hooks/asb.py +++ b/airflow/providers/microsoft/azure/hooks/asb.py @@ -348,7 +348,6 @@ def _process_message(self, msg, message_callback, receiver): message will be abandoned for future redelivery. :param receiver: The receiver that received the message. """ - print("message_callback:", message_callback) if message_callback is None: self.log.info(msg) receiver.complete_message(msg) From f9e015ad44a5fe60f640531422f5e0d537c67033 Mon Sep 17 00:00:00 2001 From: Tim Perry Date: Fri, 30 Aug 2024 08:17:56 -0700 Subject: [PATCH 3/4] Add type annotations --- .../providers/microsoft/azure/hooks/asb.py | 17 +++++++++++++--- .../microsoft/azure/hooks/test_asb.py | 20 +++++++++---------- 2 files changed, 24 insertions(+), 13 deletions(-) diff --git a/airflow/providers/microsoft/azure/hooks/asb.py b/airflow/providers/microsoft/azure/hooks/asb.py index 00d0994e65dbe..d4342d66ae2bb 100644 --- a/airflow/providers/microsoft/azure/hooks/asb.py +++ b/airflow/providers/microsoft/azure/hooks/asb.py @@ -18,7 +18,13 @@ from typing import TYPE_CHECKING, Any, Callable -from azure.servicebus import ServiceBusClient, ServiceBusMessage, ServiceBusSender +from azure.servicebus import ( + ServiceBusClient, + ServiceBusMessage, + ServiceBusReceivedMessage, + ServiceBusReceiver, + ServiceBusSender, +) from azure.servicebus.management import QueueProperties, ServiceBusAdministrationClient from airflow.hooks.base import BaseHook @@ -274,7 +280,7 @@ def send_batch_message(sender: ServiceBusSender, messages: list[str]): def receive_message( self, - queue_name, + queue_name: str, max_message_count: int | None = 1, max_wait_time: float | None = None, message_callback: MessageCallback | None = None, @@ -338,7 +344,12 @@ def receive_subscription_message( for msg in received_msgs: self._process_message(msg, message_callback, subscription_receiver) - def _process_message(self, msg, message_callback, receiver): + def _process_message( + self, + msg: ServiceBusReceivedMessage, + message_callback: MessageCallback | None, + receiver: ServiceBusReceiver, + ): """ Process the message by calling the message_callback or logging the message. diff --git a/tests/providers/microsoft/azure/hooks/test_asb.py b/tests/providers/microsoft/azure/hooks/test_asb.py index 5cbdad782cbe0..83e04833bf07f 100644 --- a/tests/providers/microsoft/azure/hooks/test_asb.py +++ b/tests/providers/microsoft/azure/hooks/test_asb.py @@ -22,7 +22,11 @@ import pytest try: - from azure.servicebus import ServiceBusClient, ServiceBusMessage, ServiceBusMessageBatch + from azure.servicebus import ( + ServiceBusClient, + ServiceBusMessage, + ServiceBusMessageBatch, + ) from azure.servicebus.management import ServiceBusAdministrationClient except ImportError: pytest.skip("Azure Service Bus not available", allow_module_level=True) @@ -266,7 +270,7 @@ def test_receive_message(self, mock_sb_client, mock_service_bus_message): ] mock_sb_client.assert_has_calls(expected_calls) - @mock.patch("azure.servicebus.ServiceBusMessage", autospec=True) + @mock.patch("azure.servicebus.ServiceBusReceivedMessage") @mock.patch(f"{MODULE}.MessageHook.get_conn", autospec=True) def test_receive_message_callback(self, mock_sb_client, mock_service_bus_message): """ @@ -326,8 +330,9 @@ def test_receive_subscription_message(self, mock_sb_client): ] mock_sb_client.assert_has_calls(expected_calls) + @mock.patch("azure.servicebus.ServiceBusReceivedMessage") @mock.patch(f"{MODULE}.MessageHook.get_conn") - def test_receive_subscription_message_callback(self, mock_sb_client): + def test_receive_subscription_message_callback(self, mock_sb_client, mock_sb_message): """ Test `receive_subscription_message` hook function and assert the function with mock value, mock the azure service bus `receive_message` function of subscription @@ -338,12 +343,9 @@ def test_receive_subscription_message_callback(self, mock_sb_client): max_wait_time = 5 hook = MessageHook(azure_service_bus_conn_id=self.conn_id) - mock_sb_message0 = ServiceBusMessage("message0") - mock_sb_message1 = ServiceBusMessage("message1") - mock_sb_client.return_value.__enter__.return_value.get_subscription_receiver.return_value.__enter__.return_value.receive_messages.return_value = [ - mock_sb_message0, - mock_sb_message1, + mock_sb_message, + mock_sb_message, ] received_messages = [] @@ -358,8 +360,6 @@ def message_callback(msg: ServiceBusMessage) -> None: ) assert len(received_messages) == 2 - assert received_messages[0] == mock_sb_message0 - assert received_messages[1] == mock_sb_message1 @pytest.mark.parametrize( "mock_subscription_name, mock_topic_name, mock_max_count, mock_wait_time", From 26c6edefa42ff827c9ced781669ad9cea1bc0f26 Mon Sep 17 00:00:00 2001 From: Tim Perry Date: Fri, 30 Aug 2024 08:18:23 -0700 Subject: [PATCH 4/4] reformat exception handling --- airflow/providers/microsoft/azure/hooks/asb.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/airflow/providers/microsoft/azure/hooks/asb.py b/airflow/providers/microsoft/azure/hooks/asb.py index d4342d66ae2bb..c90833f52fea7 100644 --- a/airflow/providers/microsoft/azure/hooks/asb.py +++ b/airflow/providers/microsoft/azure/hooks/asb.py @@ -365,8 +365,9 @@ def _process_message( else: try: message_callback(msg) - receiver.complete_message(msg) except Exception as e: self.log.error("Error processing message: %s", e) receiver.abandon_message(msg) raise e + else: + receiver.complete_message(msg)