import logging import os import re import sys import boto3 ENVIRONMENT = os.getenv('ENVIRONMENT') SERVICE_NAME = os.getenv('SERVICE_NAME') SQS_QUEUE_NAME_REGEX = os.getenv('SQS_QUEUE_NAME_REGEX') ECS_CLUSTER_NAME = os.getenv('ECS_CLUSTER_NAME') ECS_SERVICE_NAME = os.getenv('ECS_SERVICE_NAME') LOG_LEVEL = os.getenv('LOG_LEVEL', logging.INFO) AWS_DEFAULT_REGION = os.environ.get('AWS_DEFAULT_REGION', 'us-east-1') AWS_REGION = os.environ.get('AWS_REGION', AWS_DEFAULT_REGION) sqs = boto3.resource('sqs', region_name=AWS_REGION) ecs = boto3.client('ecs', region_name=AWS_REGION) cloudwatch = boto3.client('cloudwatch', region_name=AWS_REGION) def init(): if len(logging.getLogger().handlers) > 0: logging.getLogger().setLevel(LOG_LEVEL) else: logging.basicConfig( format='%(asctime)s %(levelname)s %(message)s', level=LOG_LEVEL ) logging.info('Starting init.') if not ENVIRONMENT: logging.error('You need to set "ENVIRONMENT" environment variable.') sys.exit(1) if not SERVICE_NAME: logging.error('You need to set "SERVICE_NAME" environment variable.') sys.exit(1) if not SQS_QUEUE_NAME_REGEX: logging.error('You need to set "SQS_QUEUE_NAME_REGEX" environment ' 'variable to select SQS queues to calculate message ' 'count in.') sys.exit(1) if not ECS_CLUSTER_NAME: logging.error('You need to set "ECS_CLUSTER_NAME" environment ' 'variable to get the number of ECS tasks running ' 'in the cluster.') sys.exit(1) if not ECS_SERVICE_NAME: logging.error('You need to set "ECS_SERVICE_NAME" environment ' 'variable to get the number of ECS tasks running ' 'by the service.') sys.exit(1) logging.info('Init complete.') def handler(event, context): logging.info('Starting handler.') logging.info('Getting the number of SQS messages in flight.') sqs_queue_size = get_number_of_sqs_messages_in_flight() logging.warning('The number of SQS messages in flight is %d.', sqs_queue_size) logging.info('Getting the number of ECS tasks running.') ecs_tasks_running = get_number_of_ecs_tasks_running() logging.warning('The number of ECS tasks running is %d.', ecs_tasks_running) logging.info('Publishing CloudWatch metric.') msg_per_task = round(sqs_queue_size / ecs_tasks_running) publish_cloudwath_metrics(sqs_queue_size, msg_per_task) logging.warning('The average number of SQS messages per ECS task is %d.', msg_per_task) logging.info('Handler complete.') def get_number_of_sqs_messages_in_flight(): count = 0 queue_iterator = sqs.queues.all() for queue in queue_iterator: if not re.search(SQS_QUEUE_NAME_REGEX, queue.url): continue queue_size = int(queue.attributes['ApproximateNumberOfMessages']) logging.info('Found an SQS queue "%s" with %d messages in flight.', queue.url, queue_size) count += queue_size return count def get_number_of_ecs_tasks_running(): tasks = [] paginator = ecs.get_paginator('list_tasks') response_iterator = paginator.paginate( cluster=ECS_CLUSTER_NAME, serviceName=ECS_SERVICE_NAME, desiredStatus='RUNNING', ) for page in response_iterator: tasks.extend(page['taskArns']) return len(tasks) def publish_cloudwath_metrics(sqs_queue_size, msg_per_task): namespace = '{}-{}'.format(ENVIRONMENT, SERVICE_NAME) cloudwatch.put_metric_data( Namespace=namespace, MetricData=[ { 'MetricName': 'TotalNumberOfMessagesInQueues', 'Dimensions': [ {'Name': 'QueueNameRegex', 'Value': SQS_QUEUE_NAME_REGEX} ], 'Unit': 'Count', 'Value': sqs_queue_size }, { 'MetricName': 'AverageNumberOfMessagesPerTask', 'Dimensions': [ {'Name': 'QueueNameRegex', 'Value': SQS_QUEUE_NAME_REGEX}, {'Name': 'ClusterName', 'Value': ECS_CLUSTER_NAME}, {'Name': 'ServiceName', 'Value': ECS_SERVICE_NAME} ], 'Unit': 'Count', 'Value': msg_per_task } ]) init() if __name__ == '__main__': handler(None, None)