"""Manage task worker pool.""" import json import re import backoff import requests import sentry_sdk from botocore.errorfactory import ClientError from daemon_asset_copy import config from daemon_asset_copy.connectors import logger from daemon_asset_copy.connectors import sqs as sqs_connector from daemon_asset_copy.connectors import stepfunctions from daemon_asset_copy.logic.tasks.transfer_files import transfer_from_s3_to_s3 from ddtrace import tracer def _process_task(task_message_body, sfn_client): task_token = task_message_body['task_token'] task_input = task_message_body['inputs'] or {} sfn_client.send_task_heartbeat(taskToken=task_token) outputs = transfer_from_s3_to_s3(task_input) output_validation_list = outputs if not isinstance(outputs, list): output_validation_list = [outputs] has_errors = any( re.match(r'\Aerror_.+\Z', k) for ele in output_validation_list for k, v in ele.items() ) if has_errors: sfn_client.send_task_failure( taskToken=task_token, error='error', cause=json.dumps(outputs), ) elif isinstance(task_input, list) and isinstance(outputs, list): task_input_list = task_input handler_output_list = outputs task_output_list = [] for task_input_item, handler_output_item in zip( task_input_list, handler_output_list, strict=True ): task_output_list.append({**task_input_item, **handler_output_item}) sfn_client.send_task_success( taskToken=task_token, output=json.dumps(task_output_list), ) elif isinstance(task_input, dict) and isinstance(outputs, dict): sfn_client.send_task_success( taskToken=task_token, output=json.dumps({**task_input, **outputs}), ) else: sfn_client.send_task_success( taskToken=task_token, output=json.dumps(outputs), ) def entrypoint(): """Handle an task.""" log = logger.get_current_logger() task_token = None sfn_client = stepfunctions.get_stepfunctions_client() sqs = sqs_connector.get_sqs_resource() task_queue = sqs.get_queue_by_name(QueueName=config.TASK_SQS_QUEUE_NAME) try: with tracer.trace('daemon_asset_copy_root_handler'): for _ in range(config.MAX_NUMBER_OF_SQS_MESSAGES): task_messages = task_queue.receive_messages( MaxNumberOfMessages=1, VisibilityTimeout=config.SQS_MESSAGE_VISIBILITY_TIMEOUT, WaitTimeSeconds=config.SQS_MESSAGE_WAIT_TIME_SECONDS, ) log.info(f'Grabbed {len(task_messages)} messages off the queue') for task_message in task_messages: task_message_body = json.loads(task_message.body) task_token = task_message_body['task_token'] task_message.delete() log.info( f"Inputs are :{json.dumps(task_message_body.get('inputs'))}" ) try: _process_task(task_message_body, sfn_client) except ClientError as ce: if (ce.response.get('Error') or {}).get( 'Code' ) == 'TaskTimedOut': log.info('ClientError: TaskTimedOut') else: raise ce except Exception as e: log.info('Unexpected error') _handle_unknown_error(sfn_client, task_token, e) log.info('Process complete') @tracer.wrap() def _handle_unknown_error(sfn_client, task_token, e): sentry_sdk.capture_exception() sfn_client.send_task_failure( taskToken=task_token, error=type(e).__name__, cause=str(e), ) @backoff.on_exception( backoff.expo, Exception, max_time=config.MAX_WAIT_READY_SECONDS_DATADOG ) @backoff.on_predicate(backoff.expo, max_time=config.MAX_WAIT_READY_SECONDS_DATADOG) def wait_ready_datadog(): """Wait for datadog to spin up first.""" if config.ENVIRONMENT == 'dev': return True return requests.get(config.URL_DATADOG).status_code == 404