"""Test task worker pool.""" import json import time from collections import namedtuple from unittest.mock import MagicMock, NonCallableMagicMock, call import pytest import sentry_sdk from daemon_asset_copy.connectors import sqs as sqs_connector from daemon_asset_copy.connectors import stepfunctions from daemon_asset_copy.logic import task_worker test_task_inputs = {'abc': 123} test_outputs = {'a': 'b'} test_task_token = 'task_token' test_task_name = 'test_task_name' test_worker_id = '123' test_task = { 'name': 'test_task_name', 'taskArn': 'test_task_arn', } test_exception_message = 'oops' test_message = { 'task_token': test_task_token, 'task_name': test_task_name, 'inputs': { 'abc': 123, }, } receive_messages_response = [NonCallableMagicMock(body=json.dumps(test_message))] array_test_message = { 'task_token': test_task_token, 'task_name': test_task_name, 'inputs': [ { 'abc': 123, } ], } array_receive_messages_response = [ NonCallableMagicMock(body=json.dumps(array_test_message)) ] test_task_handling_worker_parameters = namedtuple( 'test_task_handling_worker_parameters', [ 'test_description', 'receive_messages_exception', 'receive_messages_calls', 'receive_messages_response', 'handler_exception', 'handler_calls', 'send_task_success_exception', 'send_task_success_calls', 'send_task_failure_exception', 'send_task_failure_calls', 'capture_exception_calls', 'outputs', ], ) @pytest.mark.parametrize( (test_task_handling_worker_parameters._fields), [ tuple( test_task_handling_worker_parameters( test_description='Successful fetch, handle, and output result.', receive_messages_exception=None, receive_messages_calls=[ call( MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5 ) ], receive_messages_response=receive_messages_response, handler_exception=None, handler_calls=[call(test_task_inputs)], send_task_success_exception=None, send_task_success_calls=[ call( taskToken=test_task_token, output=json.dumps({**test_task_inputs, **test_outputs}), ) ], send_task_failure_exception=None, send_task_failure_calls=[], capture_exception_calls=[], outputs=test_outputs, ) ._asdict() .values() ), ( 'Successful fetch, handle, and output result with array input and output.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], array_receive_messages_response, None, [call([test_task_inputs])], None, [ call( taskToken=test_task_token, output=json.dumps([{**test_task_inputs, **test_outputs}]), ) ], None, [], [], [test_outputs], ), ( 'Successful fetch, handle, and output result with dict input' ' and array output.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [call(test_task_inputs)], None, [ call( taskToken=test_task_token, output=json.dumps([{**test_outputs}]), ) ], None, [], [], [test_outputs], ), ( 'Successful fetch, handle, and output result with errors.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [call(test_task_inputs)], None, [], None, [ call( taskToken=test_task_token, error='error', cause=json.dumps({**test_outputs, 'error_test': 'Error test.'}), ) ], [], {**test_outputs, 'error_test': 'Error test.'}, ), ( 'Successful fetch, handle, and output result with errors inside array.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [call(test_task_inputs)], None, [], None, [ call( taskToken=test_task_token, error='error', cause=json.dumps([{**test_outputs, 'error_test': 'Error test.'}]), ) ], [], [{**test_outputs, 'error_test': 'Error test.'}], ), ( 'Exception on fetch.', Exception, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [], None, [], None, [], [call()], test_outputs, ), ( 'TaskTimedOut error on send_task_heartbeat.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [], None, [], None, [], [], test_outputs, ), ( 'Non-TaskTimedOut ClientError on send_task_heartbeat.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [], None, [], None, [], [], test_outputs, ), ( 'Non-TaskTimedOut error on send_task_heartbeat.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [], None, [], None, [], [call()], test_outputs, ), ( 'Exception on handle.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, Exception(test_exception_message), [call(test_task_inputs)], None, [], None, [ call( taskToken=test_task_token, error=type(Exception(test_exception_message)).__name__, cause=str(Exception(test_exception_message)), ) ], [call()], test_outputs, ), ( 'Exception on send_task_success.', None, [call(MaxNumberOfMessages=1, VisibilityTimeout=120, WaitTimeSeconds=5)], receive_messages_response, None, [call(test_task_inputs)], Exception(test_exception_message), [ call( taskToken=test_task_token, output=json.dumps({**test_task_inputs, **test_outputs}), ) ], None, [ call( taskToken=test_task_token, error=type(Exception(test_exception_message)).__name__, cause=str(Exception(test_exception_message)), ) ], [call()], test_outputs, ), ], ) def test_task_handling_worker( mocker, test_description, receive_messages_exception, receive_messages_calls, receive_messages_response, handler_exception, handler_calls, send_task_success_exception, send_task_success_calls, send_task_failure_exception, send_task_failure_calls, capture_exception_calls, outputs, ): """Test task_handling_worker.""" mocker.patch.object(stepfunctions, 'get_stepfunctions_client', autospec=True) mocker.patch.object(sentry_sdk, 'capture_exception', autospec=True) mocker.patch.object(sqs_connector, 'get_sqs_resource', autospec=True) ( # fmt: off sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .receive_messages .return_value # fmt: on ) = receive_messages_response ( # fmt: off sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .receive_messages .side_effect # fmt: on ) = receive_messages_exception ( # fmt: off stepfunctions .get_stepfunctions_client .return_value .send_task_success .side_effect # fmt: on ) = send_task_success_exception ( # fmt: off stepfunctions .get_stepfunctions_client .return_value .send_task_failure .side_effect # fmt: on ) = send_task_failure_exception mock_handler = MagicMock(return_value=outputs, side_effect=handler_exception) mocker.patch.object(task_worker, 'transfer_from_s3_to_s3', mock_handler) mocker.patch.object(time, 'sleep', autospec=True) task_worker.entrypoint() ( # fmt: off sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .receive_messages .assert_has_calls( receive_messages_calls ) # fmt: on ) ( # fmt: off stepfunctions .get_stepfunctions_client .return_value .send_task_success .assert_has_calls( send_task_success_calls ) # fmt: on ) ( # fmt: off stepfunctions .get_stepfunctions_client .return_value .send_task_failure .assert_has_calls( send_task_failure_calls ) # fmt: on ) (mock_handler.asset_has_calls(handler_calls))