"""Manage activity task worker pool.""" import concurrent.futures import json import random import re import threading import time import uuid from botocore.errorfactory import ClientError import sentry_sdk 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.constants import exceptions from video.logic import activity_task_worker_throttlers from video.utils import exception def is_fifo(queue_url): """Determine queue type.""" return queue_url.endswith('fifo') def merge_if_multiple_inputs(inputs): """If inputs is an array of objects merge them into one. Args: inputs (list): Inputs. This is the case if the previous step in the state machine was a parallel step. Returns: dict: Merged inputs. """ if isinstance(inputs, dict): return inputs return { k: v for inputs_item in inputs for k, v in merge_if_multiple_inputs(inputs_item).items()} def _handle_unexpected_polling_exception(sfn_client, activity_task_token, e): """Handle unexpected polling exception.""" sentry_sdk.capture_exception() try: sfn_client.send_task_failure( taskToken=activity_task_token, error=type(e).__name__, cause=str(e), ) except Exception: sentry_sdk.capture_exception() def activity_task_polling_worker(worker_id, activity): """Poll for activity task. Args: worker_id: some unique id identifying this activity_task_polling_worker activity: activity object with arn """ sfn_client = stepfunctions.get_stepfunctions_client() activity_task_token = None # Get the service resource sqs = sqs_connector.get_sqs_resource() # Get the queue activity_task_queue = sqs.get_queue_by_name( QueueName=config.ACTIVITY_TASK_SQS_QUEUE_NAME) time.sleep(60 * random.random()) while activity_task_worker_throttlers.is_polling_worker_activated(): try: activity_arn = activity['activityArn'] activity_task = sfn_client.get_activity_task( activityArn=activity_arn, workerName=worker_id, ) activity_task_token = activity_task.get('taskToken') # If there are no activity tasks within 60 seconds the request # times out and we have to make another one. if not activity_task_token: continue activity_task_message = { 'activity_task': activity_task, 'activity_name': activity['name'], } if is_fifo(config.ACTIVITY_TASK_SQS_QUEUE_NAME): activity_task_queue.send_message( MessageBody=json.dumps(activity_task_message), MessageGroupId=worker_id, ) else: activity_task_queue.send_message( MessageBody=json.dumps(activity_task_message), ) except ClientError as ce: if (ce.response.get('Error') or {}).get( 'Code') == 'ThrottlingException': time.sleep( config.MAX_TIMEOUT_AFTER_AWS_THROTTLING_EXCEPTION_SECONDS * random.random()) # noqa else: _handle_unexpected_polling_exception( sfn_client, activity_task_token, ce) except Exception as e: _handle_unexpected_polling_exception( sfn_client, activity_task_token, e) def send_task_heartbeat( activity_task_token=None, heartbeat_timer_container=None, ): """ Send heartbeat for activity task. Args: activity_task_token: Activity task token. heartbeat_timer_container: Container to put timer in. """ sfn_client = stepfunctions.get_stepfunctions_client() try: sfn_client.send_task_heartbeat(taskToken=activity_task_token) timer = threading.Timer( interval=config.HEARTBEAT_PERIOD_SECONDS, function=send_task_heartbeat, args=[], kwargs={ 'activity_task_token': activity_task_token, 'heartbeat_timer_container': heartbeat_timer_container, }, ) timer.start() heartbeat_timer_container['timer'] = timer except ClientError as ce: if (ce.response.get('Error') or {}).get('Code') != 'TaskTimedOut': _handle_unknown_error(sfn_client, activity_task_token, ce) exception.reraise_exception(ce) except Exception as e: _handle_unknown_error(sfn_client, activity_task_token, e) exception.reraise_exception(e) def _handle_unknown_error(sfn_client, activity_task_token, e): sentry_sdk.capture_exception() sfn_client.send_task_failure( taskToken=activity_task_token, error=type(e).__name__, cause=str(e), ) def activity_task_handling_worker(): """Handle an activity task.""" sfn_client = stepfunctions.get_stepfunctions_client() activity_task_input = None activity_task_token = None activity_task_message = None # Get the service resource sqs = sqs_connector.get_sqs_resource() heartbeat_timer_container = {} # Get the queue activity_task_queue = sqs.get_queue_by_name( QueueName=config.ACTIVITY_TASK_SQS_QUEUE_NAME) while activity_task_worker_throttlers.is_handling_worker_activated(): try: activity_task_messages = ( activity_task_queue.receive_messages(MaxNumberOfMessages=1)) if not activity_task_messages: continue activity_task_message, = activity_task_messages activity_task_message_body = ( json.loads(activity_task_message.body)) activity_task_message.delete() activity_name = activity_task_message_body['activity_name'] activity_task = activity_task_message_body['activity_task'] activity_task_inputs = json.loads( activity_task.get('input') or '{}') activity_task_token = activity_task.get('taskToken') activity_task_input = merge_if_multiple_inputs( activity_task_inputs) handler = activity_constants.ACTIVITIES_TO_HANDLERS[activity_name] send_task_heartbeat( activity_task_token=activity_task_token, heartbeat_timer_container=heartbeat_timer_container, ) activity_task_worker_throttlers.handling_worker_throttler() outputs = handler(activity_task_input) has_errors = any( re.match(r'\Aerror_.+\Z', k) for k, v in outputs.items()) if has_errors: sfn_client.send_task_failure( taskToken=activity_task_token, error='error', cause=json.dumps(outputs), ) else: sfn_client.send_task_success( taskToken=activity_task_token, output=json.dumps({**activity_task_input, **outputs}), ) except exceptions.NotCurrentlyAcceptingWork: pass except ClientError as ce: if (ce.response.get('Error') or {}).get('Code') != 'TaskTimedOut': _handle_unknown_error(sfn_client, activity_task_token, ce) except Exception as e: _handle_unknown_error(sfn_client, activity_task_token, e) finally: if heartbeat_timer_container: heartbeat_timer_container['timer'].cancel() heartbeat_timer_container = {} def start(): """Start worker pool.""" activities_to_poll_for_tasks = stepfunctions.list_activities() num_pollers = len(activities_to_poll_for_tasks) num_handlers = config.NUM_ACTIVITY_TASK_HANDLING_WORKERS with concurrent.futures.ProcessPoolExecutor( max_workers=num_pollers + num_handlers ) as executor: for _ in range(num_handlers): executor.submit( activity_task_handling_worker, ) for activity in activities_to_poll_for_tasks: executor.submit( activity_task_polling_worker, '{}{}-{}-{}'.format( config.WORKER_NAME_PREFIX, config.ENVIRONMENT, config.SERVICE_NAME, str(uuid.uuid1()) ), activity )