From b250f98e463efd9876c2d40094b16c06113bfa72 Mon Sep 17 00:00:00 2001 From: dbarrundia3 Date: Tue, 19 Jul 2022 16:20:27 -0600 Subject: [PATCH] SQSPublishOperator should allow sending messages to a FIFO Queue --- airflow/providers/amazon/aws/hooks/sqs.py | 20 +++++++---- airflow/providers/amazon/aws/operators/sqs.py | 16 +++++++-- .../amazon/aws/operators/test_sqs.py | 33 +++++++++++++++++++ 3 files changed, 61 insertions(+), 8 deletions(-) diff --git a/airflow/providers/amazon/aws/hooks/sqs.py b/airflow/providers/amazon/aws/hooks/sqs.py index 97fdff01d1d3a..bb47dbb2533db 100644 --- a/airflow/providers/amazon/aws/hooks/sqs.py +++ b/airflow/providers/amazon/aws/hooks/sqs.py @@ -59,6 +59,7 @@ def send_message( message_body: str, delay_seconds: int = 0, message_attributes: Optional[Dict] = None, + message_group_id: str = None, ) -> Dict: """ Send message to the queue @@ -72,14 +73,21 @@ def send_message( :param message_attributes: additional attributes for the message (default: None) For details of the attributes parameter see :py:meth:`botocore.client.SQS.send_message` :type message_attributes: dict + :param message_group_id: This parameter applies only to FIFO (first-in-first-out) queues. (default: None) + For details of the attributes parameter see :py:meth:`botocore.client.SQS.send_message` + :type message_group_id: str :return: dict with the information about the message sent For details of the returned value see :py:meth:`botocore.client.SQS.send_message` :rtype: dict """ - return self.get_conn().send_message( - QueueUrl=queue_url, - MessageBody=message_body, - DelaySeconds=delay_seconds, - MessageAttributes=message_attributes or {}, - ) + params = { + 'QueueUrl': queue_url, + 'MessageBody': message_body, + 'DelaySeconds': delay_seconds, + 'MessageAttributes': message_attributes or {}, + } + if message_group_id: + params['MessageGroupId'] = message_group_id + + return self.get_conn().send_message(**params) diff --git a/airflow/providers/amazon/aws/operators/sqs.py b/airflow/providers/amazon/aws/operators/sqs.py index ae113cf250724..1560f24793f32 100644 --- a/airflow/providers/amazon/aws/operators/sqs.py +++ b/airflow/providers/amazon/aws/operators/sqs.py @@ -39,11 +39,20 @@ class SQSPublishOperator(BaseOperator): :type message_attributes: dict :param delay_seconds: message delay (templated) (default: 1 second) :type delay_seconds: int + :param message_group_id: This parameter applies only to FIFO (first-in-first-out) queues. (default: None) + For details of the attributes parameter see :py:meth:`botocore.client.SQS.send_message` + :type message_group_id: str :param aws_conn_id: AWS connection id (default: aws_default) :type aws_conn_id: str """ - template_fields = ('sqs_queue', 'message_content', 'delay_seconds', 'message_attributes') + template_fields = ( + 'sqs_queue', + 'message_content', + 'delay_seconds', + 'message_attributes', + 'message_group_id', + ) template_fields_renderers = {'message_attributes': 'json'} ui_color = '#6ad3fa' @@ -54,6 +63,7 @@ def __init__( message_content: str, message_attributes: Optional[dict] = None, delay_seconds: int = 0, + message_group_id: str = None, aws_conn_id: str = 'aws_default', **kwargs, ): @@ -63,6 +73,7 @@ def __init__( self.message_content = message_content self.delay_seconds = delay_seconds self.message_attributes = message_attributes or {} + self.message_group_id = message_group_id def execute(self, context): """ @@ -81,8 +92,9 @@ def execute(self, context): message_body=self.message_content, delay_seconds=self.delay_seconds, message_attributes=self.message_attributes, + message_group_id=self.message_group_id, ) - self.log.info('result is send_message is %s', result) + self.log.info('result of send_message is %s', result) return result diff --git a/tests/providers/amazon/aws/operators/test_sqs.py b/tests/providers/amazon/aws/operators/test_sqs.py index ea695194aca4d..8a9a8596098df 100644 --- a/tests/providers/amazon/aws/operators/test_sqs.py +++ b/tests/providers/amazon/aws/operators/test_sqs.py @@ -20,6 +20,8 @@ import unittest from unittest.mock import MagicMock +import pytest +from botocore.exceptions import ClientError from moto import mock_sqs from airflow.models.dag import DAG @@ -32,6 +34,9 @@ QUEUE_NAME = 'test-queue' QUEUE_URL = f'https://{QUEUE_NAME}' +FIFO_QUEUE_NAME = 'test-queue.fifo' +FIFO_QUEUE_URL = f'https://{FIFO_QUEUE_NAME}' + class TestSQSPublishOperator(unittest.TestCase): def setUp(self): @@ -66,3 +71,31 @@ def test_execute_success(self): context_calls = [] assert self.mock_context['ti'].method_calls == context_calls, "context call should be same" + + @mock_sqs + def test_execute_failure_fifo_queue(self): + self.operator.sqs_queue = FIFO_QUEUE_URL + self.sqs_hook.create_queue(FIFO_QUEUE_NAME, attributes={'FifoQueue': 'true'}) + with pytest.raises(ClientError) as ctx: + self.operator.execute(self.mock_context) + err_msg = ( + "An error occurred (MissingParameter) when calling the SendMessage operation: The request must " + "contain the parameter MessageGroupId." + ) + assert err_msg == str(ctx.value) + + @mock_sqs + def test_execute_success_fifo_queue(self): + self.operator.sqs_queue = FIFO_QUEUE_URL + self.operator.message_group_id = "abc" + self.sqs_hook.create_queue(FIFO_QUEUE_NAME, attributes={'FifoQueue': 'true'}) + result = self.operator.execute(self.mock_context) + assert 'MD5OfMessageBody' in result + assert 'MessageId' in result + message = self.sqs_hook.get_conn().receive_message( + QueueUrl=FIFO_QUEUE_URL, AttributeNames=['MessageGroupId'] + ) + assert len(message['Messages']) == 1 + assert message['Messages'][0]['MessageId'] == result['MessageId'] + assert message['Messages'][0]['Body'] == 'hello' + assert message['Messages'][0]['Attributes']['MessageGroupId'] == 'abc'