"""Wrapper for bulk rollback service.""" import logging import multiprocessing import os import sys import boto3 from botocore.exceptions import ClientError from dotenv import load_dotenv import deploy # Load env file if it exists load_dotenv(verbose=True) logger = logging.getLogger() stream_handler = logging.StreamHandler(sys.stdout) stream_handler.setLevel('INFO') logger.handlers = [stream_handler] ENVIRONMENT = os.environ.get('Environment', 'dev') AWS_REGION = os.environ.get('AWS_REGION', 'us-east-1') CLUSTER_NAME = os.environ.get('CLUSTER_NAME') SERVICE_NAME_PREFIXES = os.environ.get('SERVICE_NAME_PREFIXES').split(',') def get_services(client, enviromment, service_name_prefix): """ Get list of ECS services given a prefix. Returns: list(dict): list of ECS service dictionaries """ matching_services = [] services = client.list_services( cluster=CLUSTER_NAME, maxResults=100, )['serviceArns'] # Get all services and find the matching ones. [matching_services.append(service_arn.split('/')[-1]) for service_arn in services if f'{enviromment}-{service_name_prefix}' in service_arn] # noqa logger.info(f'Found matching services: {matching_services}') return matching_services def get_service_task_definition(client, service): """ Get current task definition for a service. Returns: str: Task definition revision """ latest_task = client.describe_task_definition( taskDefinition=service) current_task_def = latest_task['taskDefinition'] task_def_revision = current_task_def['taskDefinitionArn'].split(':')[-1] return task_def_revision def revert(full_service_name): client = boto3.client('ecs', region_name=AWS_REGION) logger.info(f'Running rollback for {full_service_name}') try: task_def_revision = get_service_task_definition( client, full_service_name) revert_revision = str(int(task_def_revision) - 1) # Strip the environment prefix to get the service name only service_name = full_service_name.lstrip(f'{ENVIRONMENT}-') logger.info(f'Revert task revision is {revert_revision}') deploy.deploy_mode = 'REVERT' deploy.cluster_name = CLUSTER_NAME deploy.service_name = service_name deploy.fargate_service_name = full_service_name deploy.task_family = full_service_name deploy.container_name = service_name deploy.update_timeout = 600 deploy.revert_task_definition_number = revert_revision deploy.verify_mode = 'TASK_RUNNING' deploy.main() except ClientError as error: # We want to continue to the next service, so just log the error logger.error(f'Error rolling back {full_service_name}: {error}') def main(): """Main entry point function.""" client = boto3.client('ecs', region_name=AWS_REGION) for prefix in SERVICE_NAME_PREFIXES: logger.info(f'Running rollback for prefix {prefix}') service_list = get_services(client, ENVIRONMENT, prefix) with multiprocessing.Pool(processes=len(service_list)) as pool: _result = pool.map(revert, sorted(service_list)) if __name__ == '__main__': main()