"""Test activity task worker pool.""" from collections import namedtuple import concurrent.futures import json import threading import time from unittest.mock import call from unittest.mock import MagicMock from unittest.mock import NonCallableMagicMock import uuid from botocore.errorfactory import ClientError from botocore.vendored.requests.exceptions import ReadTimeout import pytest import sentry_sdk from tests.test_utils import match_any from video import config from video.connectors import sqs as sqs_connector from video.connectors import stepfunctions from video.constants import activities as activity_constants from video.logic import activity_task_worker_pool from video.logic import activity_task_worker_throttlers from video.utils import exception def test_start(mocker): """Test start.""" activities_to_poll_for_tasks = [ {'activityArn': 1}, {'activityArn': 1}, {'activityArn': 1}] mocker.patch.object( stepfunctions, 'list_activities', return_value=activities_to_poll_for_tasks, autospec=True) mocker.patch.object( activity_task_worker_pool, 'activity_task_handling_worker', autospec=True) mocker.patch.object( concurrent.futures, 'ProcessPoolExecutor', autospec=True) mocker.patch.object(uuid, 'uuid1', autospec=True) mocker.patch.object(time, 'sleep', autospec=True) num_pollers = len(activities_to_poll_for_tasks) num_handlers = config.NUM_ACTIVITY_TASK_HANDLING_WORKERS activity_task_worker_pool.start() concurrent.futures.ProcessPoolExecutor.assert_called_once_with( max_workers=num_pollers + num_handlers) stepfunctions.list_activities.assert_called_once() (concurrent .futures .ProcessPoolExecutor .return_value .__enter__ .return_value .submit .assert_has_calls([ *[ call( activity_task_worker_pool.activity_task_handling_worker, ) for _ in range(num_handlers) ], *[ call( activity_task_worker_pool.activity_task_polling_worker, '{}{}-{}-{}'.format( config.WORKER_NAME_PREFIX, config.ENVIRONMENT, config.SERVICE_NAME, str(uuid.uuid1.return_value) ), activity, ) for activity in activities_to_poll_for_tasks ], ])) test_activity_task_inputs = {'abc': 123} test_outputs = {'a': 'b'} test_task_token = 'applesauce_bananas' test_activity_name = 'test_activity_name' test_worker_id = '123' test_activity = { 'name': 'test_activity_name', 'activityArn': 'test_activity_arn', } test_exception_message = 'oops' test_message = { 'activity_task': { 'input': json.dumps(test_activity_task_inputs), 'taskToken': test_task_token, }, 'activity_name': test_activity_name } receive_messages_response = [NonCallableMagicMock( body=json.dumps(test_message) )] test_activity_task_handling_worker_parameters = namedtuple( 'test_activity_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', 'send_task_heartbeat_exception', 'send_task_heartbeat_calls', 'threading_timer_calls', 'timer_start_call_count', 'timer_cancel_call_count', 'capture_exception_calls', 'outputs', ] ) threading_timer_mock = MagicMock() @pytest.mark.parametrize(( test_activity_task_handling_worker_parameters._fields ), [ tuple(test_activity_task_handling_worker_parameters( test_description='Successful fetch, handle, and output result.', receive_messages_exception=None, receive_messages_calls=[call(MaxNumberOfMessages=1)], receive_messages_response=receive_messages_response, handler_exception=None, handler_calls=[call(test_activity_task_inputs)], send_task_success_exception=None, send_task_success_calls=[call( taskToken=test_task_token, output=json.dumps( {**test_activity_task_inputs, **test_outputs}))], send_task_failure_exception=None, send_task_failure_calls=[], send_task_heartbeat_exception=None, send_task_heartbeat_calls=[call(taskToken='applesauce_bananas')], threading_timer_calls=[call( interval=30, function=activity_task_worker_pool.send_task_heartbeat, args=[], kwargs={ 'activity_task_token': test_task_token, 'heartbeat_timer_container': { 'timer': threading_timer_mock.return_value, }, }, )], timer_start_call_count=1, timer_cancel_call_count=1, capture_exception_calls=[], outputs=test_outputs, )._asdict().values()), ( 'Successful fetch, handle, and output result with array input.', None, [call(MaxNumberOfMessages=1)], receive_messages_response, None, [call(test_activity_task_inputs)], None, [call( taskToken=test_task_token, output=json.dumps({**test_activity_task_inputs, **test_outputs}))], None, [], None, [call(taskToken='applesauce_bananas')], [call( interval=30, function=activity_task_worker_pool.send_task_heartbeat, args=[], kwargs={ 'activity_task_token': test_task_token, 'heartbeat_timer_container': { 'timer': threading_timer_mock.return_value, }, }, )], 1, 1, [], test_outputs, ), ( 'Successful fetch, handle, and output result with errors.', None, [call(MaxNumberOfMessages=1)], receive_messages_response, None, [call(test_activity_task_inputs)], None, [], None, [call( taskToken=test_task_token, error='error', cause=json.dumps({**test_outputs, 'error_test': 'Error test.'}), )], None, [call(taskToken='applesauce_bananas')], [call( interval=30, function=activity_task_worker_pool.send_task_heartbeat, args=[], kwargs={ 'activity_task_token': test_task_token, 'heartbeat_timer_container': { 'timer': threading_timer_mock.return_value, }, }, )], 1, 1, [], {**test_outputs, 'error_test': 'Error test.'}, ), ( 'Long Poll Timeout.', None, [call(MaxNumberOfMessages=1)], [], None, [], None, [], None, [], None, [], [], 0, 0, [], test_outputs, ), ( 'Exception on fetch.', Exception, [call(MaxNumberOfMessages=1)], receive_messages_response, None, [], None, [], None, [], None, [], [], 0, 0, [call()], test_outputs, ), ( 'TaskTimedOut error on send_task_heartbeat.', None, [call(MaxNumberOfMessages=1)], receive_messages_response, None, [], None, [], None, [], ClientError({'Error': {'Code': 'TaskTimedOut'}}, 'asdf'), [], [], 0, 0, [], test_outputs, ), ( 'Non-TaskTimedOut ClientError on send_task_heartbeat.', None, [call(MaxNumberOfMessages=1)], receive_messages_response, None, [], None, [], None, [], ClientError({'Error': {'Code': 'NotTaskTimedOutWhateves'}}, 'asdf'), [], [], 0, 0, [], test_outputs, ), ( 'Non-TaskTimedOut error on send_task_heartbeat.', None, [call(MaxNumberOfMessages=1)], receive_messages_response, None, [], None, [], None, [], ClientError({}, 'asdf'), [], [], 0, 0, [call()], test_outputs, ), ( 'Exception on handle.', None, [call(MaxNumberOfMessages=1)], receive_messages_response, Exception(test_exception_message), [call(test_activity_task_inputs)], None, [], None, [call( taskToken=test_task_token, error=type(Exception(test_exception_message)).__name__, cause=str(Exception(test_exception_message)), )], None, [call(taskToken='applesauce_bananas')], [call( interval=30, function=activity_task_worker_pool.send_task_heartbeat, args=[], kwargs={ 'activity_task_token': test_task_token, 'heartbeat_timer_container': { 'timer': threading_timer_mock.return_value, }, }, )], 1, 1, [call()], test_outputs, ), ( 'Exception on send_task_success.', None, [call(MaxNumberOfMessages=1)], receive_messages_response, None, [call(test_activity_task_inputs)], Exception(test_exception_message), [call( taskToken=test_task_token, output=json.dumps({**test_activity_task_inputs, **test_outputs}))], None, [call( taskToken=test_task_token, error=type(Exception(test_exception_message)).__name__, cause=str(Exception(test_exception_message)), )], None, [call(taskToken='applesauce_bananas')], [call( interval=30, function=activity_task_worker_pool.send_task_heartbeat, args=[], kwargs={ 'activity_task_token': test_task_token, 'heartbeat_timer_container': { 'timer': threading_timer_mock.return_value, }, }, )], 1, 1, [call()], test_outputs, ), ]) def test_activity_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, send_task_heartbeat_exception, send_task_heartbeat_calls, threading_timer_calls, timer_start_call_count, timer_cancel_call_count, capture_exception_calls, outputs, ): """Test activity_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) (sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .receive_messages .return_value) = receive_messages_response (sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .receive_messages .side_effect) = receive_messages_exception (stepfunctions .get_stepfunctions_client .return_value .send_task_heartbeat .side_effect) = send_task_heartbeat_exception (stepfunctions .get_stepfunctions_client .return_value .send_task_success .side_effect) = send_task_success_exception (stepfunctions .get_stepfunctions_client .return_value .send_task_failure .side_effect) = send_task_failure_exception mocker.patch.object( activity_task_worker_throttlers, 'is_handling_worker_activated', autospec=True, side_effect=[True, False]) mocker.patch.object( activity_task_worker_throttlers, 'handling_worker_throttler', autospec=True, return_value=True) mocker.patch.dict( activity_constants .ACTIVITIES_TO_HANDLERS, {test_activity_name: MagicMock( return_value=outputs, side_effect=handler_exception)}) mocker.patch.object(time, 'sleep', autospec=True) threading_timer_mock.reset_mock() mocker.patch.object( threading, 'Timer', return_value=threading_timer_mock.return_value, autospec=True) activity_task_worker_pool.activity_task_handling_worker() assert ( activity_task_worker_throttlers .is_handling_worker_activated.call_count) == 2 (sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .receive_messages.assert_has_calls(receive_messages_calls)) (stepfunctions .get_stepfunctions_client .return_value .send_task_success .assert_has_calls(send_task_success_calls)) (stepfunctions .get_stepfunctions_client .return_value .send_task_failure .assert_has_calls(send_task_failure_calls)) (stepfunctions .get_stepfunctions_client .return_value .send_task_heartbeat .assert_has_calls(send_task_heartbeat_calls)) threading.Timer.assert_has_calls(threading_timer_calls) assert threading.Timer.return_value.start.call_count == ( timer_start_call_count) assert threading.Timer.return_value.cancel.call_count == ( timer_cancel_call_count) sentry_sdk.capture_exception.assert_has_calls(capture_exception_calls) (activity_constants .ACTIVITIES_TO_HANDLERS[test_activity_name] .assert_has_calls(handler_calls)) @pytest.mark.parametrize(( 'test_description', 'get_activity_task_exception', 'get_activity_task_calls', 'send_message_exception', 'send_message_calls', 'send_task_failure_exception', 'send_task_failure_calls', 'capture_exception_calls', 'inputs', 'outputs', 'task_token', 'expected_sleep_calls', ), [ ( 'Successful fetch, handle, and output result.', None, [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], None, [call(MessageBody=json.dumps(test_message))], None, [], [], json.dumps([test_activity_task_inputs, {'ok': 'cool'}]), test_outputs, test_task_token, [call(match_any(float))], ), ( 'Successful fetch, handle, and output result with array input.', None, [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], None, [call(MessageBody=json.dumps(test_message))], None, [], [], json.dumps(test_activity_task_inputs), test_outputs, test_task_token, [call(match_any(float))], ), ( 'Long Poll Timeout.', ReadTimeout, [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], None, [], None, [], [], None, test_outputs, None, [call(match_any(float))], ), ( 'Exception on fetch.', Exception, [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], None, [], None, [], [call()], json.dumps(test_activity_task_inputs), test_outputs, test_task_token, [call(match_any(float))], ), ( 'ThrottlingException on fetch.', ClientError({'Error': {'Code': 'ThrottlingException'}}, 'asdf'), [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], None, [], None, [], [], json.dumps(test_activity_task_inputs), test_outputs, test_task_token, [call(match_any(float)), call(match_any(float))], ), ( 'NotAThrottlingException on fetch.', ClientError({'Error': {'Code': 'NotAThrottlingException'}}, 'asdf'), [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], None, [], None, [], [call()], json.dumps(test_activity_task_inputs), test_outputs, test_task_token, [call(match_any(float))], ), ( 'Exception on handle.', None, [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], Exception(test_exception_message), [call(MessageBody=json.dumps(test_message))], None, [call( taskToken=test_task_token, error=type(Exception(test_exception_message)).__name__, cause=str(Exception(test_exception_message)), )], [call()], json.dumps(test_activity_task_inputs), test_outputs, test_task_token, [call(match_any(float))], ), ( 'Exception on send_task_failure.', None, [call( activityArn=test_activity['activityArn'], workerName=test_worker_id)], Exception(test_exception_message), [call(MessageBody=json.dumps(test_message))], Exception, [call( taskToken=test_task_token, error=type(Exception(test_exception_message)).__name__, cause=str(Exception(test_exception_message)), )], [call(), call()], json.dumps(test_activity_task_inputs), test_outputs, test_task_token, [call(match_any(float))], ), ]) def test_activity_task_polling_worker( mocker, test_description, get_activity_task_exception, get_activity_task_calls, send_message_exception, send_message_calls, send_task_failure_exception, send_task_failure_calls, capture_exception_calls, inputs, outputs, task_token, expected_sleep_calls, ): """Test activity_task_polling_worker.""" mocker.patch.object( config, 'ACTIVITY_TASK_SQS_QUEUE_NAME', 'https://queue.amazonaws.com/111111/test-queue') mocker.patch.object( stepfunctions, 'get_stepfunctions_client', autospec=True) mocker.patch.object( sentry_sdk, 'capture_exception', autospec=True) (stepfunctions .get_stepfunctions_client .return_value .get_activity_task .return_value) = { 'input': json.dumps(test_activity_task_inputs), 'taskToken': test_task_token, } (stepfunctions .get_stepfunctions_client .return_value .get_activity_task .side_effect) = get_activity_task_exception (stepfunctions .get_stepfunctions_client .return_value .send_task_failure .side_effect) = send_task_failure_exception mocker.patch.object( activity_task_worker_throttlers, 'is_polling_worker_activated', autospec=True, side_effect=[True, False]) mocker.patch.object(sqs_connector, 'get_sqs_resource', autospec=True) (sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .send_message .side_effect) = send_message_exception mocker.patch.object(time, 'sleep') activity_task_worker_pool.activity_task_polling_worker( test_worker_id, test_activity) assert ( activity_task_worker_throttlers .is_polling_worker_activated.call_count) == 2 (stepfunctions .get_stepfunctions_client .return_value .get_activity_task .assert_has_calls(get_activity_task_calls)) (stepfunctions .get_stepfunctions_client .return_value .send_task_failure .assert_has_calls(send_task_failure_calls)) sentry_sdk.capture_exception.assert_has_calls(capture_exception_calls) (sqs_connector .get_sqs_resource .return_value .get_queue_by_name .return_value .send_message.assert_has_calls(send_message_calls)) time.sleep.assert_has_calls(expected_sleep_calls) test_exception = Exception(test_exception_message) test_client_error = ClientError({'Error': {'Code': 'asdfasdf'}}, 'asdf') test_timeout_client_error = ClientError( {'Error': {'Code': 'TaskTimedOut'}}, 'asdf') @pytest.mark.parametrize(( 'send_task_heartbeat_exception', 'reraise_exception_calls', 'capture_exception_calls', 'send_task_failure_calls', 'threading_timer_calls', 'expected_heartbeat_timer_container', ), [ ( test_exception, [call(test_exception)], [call()], [call( taskToken=test_task_token, error=type(test_exception).__name__, cause=str(test_exception), )], [], {}, ), ( test_timeout_client_error, [call(test_timeout_client_error)], [], [], [], {}, ), ( test_client_error, [call(test_client_error)], [call()], [call( taskToken=test_task_token, error=type(test_client_error).__name__, cause=str(test_client_error), )], [], {}, ), ( None, [], [], [], [call( interval=30, function=activity_task_worker_pool.send_task_heartbeat, args=[], kwargs={ 'activity_task_token': test_task_token, 'heartbeat_timer_container': { 'timer': threading_timer_mock.return_value, }, }, )], {'timer': threading_timer_mock.return_value}, ), ]) def test_send_task_heartbeat( mocker, send_task_heartbeat_exception, reraise_exception_calls, capture_exception_calls, send_task_failure_calls, threading_timer_calls, expected_heartbeat_timer_container, ): """Test send_task_heartbeat.""" mocker.patch.object( stepfunctions, 'get_stepfunctions_client', autospec=True) mocker.patch.object(exception, 'reraise_exception', autospec=True) mocker.patch.object(sentry_sdk, 'capture_exception', autospec=True) mocker.patch.object( threading, 'Timer', return_value=threading_timer_mock.return_value, autospec=True) (stepfunctions .get_stepfunctions_client .return_value .send_task_heartbeat .side_effect) = send_task_heartbeat_exception heartbeat_timer_container = {} activity_task_worker_pool.send_task_heartbeat( test_task_token, heartbeat_timer_container) assert exception.reraise_exception.mock_calls == reraise_exception_calls assert sentry_sdk.capture_exception.mock_calls == capture_exception_calls assert ( stepfunctions .get_stepfunctions_client .return_value .send_task_failure .mock_calls) == send_task_failure_calls assert ( stepfunctions .get_stepfunctions_client .return_value .send_task_heartbeat .mock_calls) == [call(taskToken=test_task_token)] assert threading.Timer.mock_calls == threading_timer_calls assert heartbeat_timer_container == expected_heartbeat_timer_container