import dataclasses import datetime import functools import logging import time import typing import uuid import boto3 # import mysql.connector.connection import psycopg2 import pymysql from psycopg2 import sql import psycopg2.extensions logger = logging.getLogger(__name__) @dataclasses.dataclass class Credentials: username: str password: str @classmethod def new(cls): username_format = "u{}" # must start with a letter return cls( username=username_format.format(uuid.uuid4().hex), password=uuid.uuid4().hex, ) # make sure it starts with a letter @dataclasses.dataclass class DbConnection: endpoint: str port: int @dataclasses.dataclass class DbInstance(DbConnection): cluster_id: str arn: str reader_endpoint: typing.Optional[str] = None database_name: typing.Optional[str] = None user: typing.Optional[Credentials] = None master_user: typing.Optional[Credentials] = None def retry(timeout: int, backoff_seconds: int = 10, bubble_errors: typing.List[typing.Type[Exception]] = None): """ Tries calling a function {timeout} seconds until it succeeds or gives up and throw an error :param timeout: Timeout in seconds :param backoff_seconds: Wait time between retries :param bubble_errors: List of exceptions to bubble up :return: Wrapped function """ bubble_errors = tuple(bubble_errors or []) def wrapped(func: typing.Callable): @functools.wraps(func) def func_retrier(*args, **kwargs): waited = 0 started_at = datetime.datetime.now() while True: try: logger.debug(f"Calling {func}") result = func(*args, **kwargs) logger.debug(f"{func} completed in {(datetime.datetime.now() - started_at)}") return result except KeyboardInterrupt: raise except tuple(bubble_errors): raise except Exception as e: doze = min(timeout - waited, backoff_seconds) times = (timeout - waited) // backoff_seconds logger.debug(f"{func} failed, will try {times} times in {doze} sec", exc_info=True) time.sleep(doze) waited += backoff_seconds if waited >= timeout: raise TimeoutError(f"Couldn't get a successful result from {func} in {timeout} seconds") from e return func_retrier return wrapped def paginate_by_marker( func: callable, list_field: str, ) -> typing.Iterable: """ Simplifies paginating over AWS results with a simpler interface than Paginators. :param func: A function that accepts a parameter named `Marker` :param list_field: Result field to iterate on """ marker = "" while True: result = func(Marker=marker) items = result.get(list_field, []) yield from items if "Marker" not in result: break marker = result["Marker"] class PostgresqlService: db_type = "Postgresql" db_engine = "aurora-postgresql" db_engine_version = "10.11" default_master_database_name = "postgres" cluster_id_format = "pg-{}" # must start with a letter tenant_database_name_format = "db{}" # must start with a letter def __init__(self, region: str): self.rds = boto3.client("rds", region_name=region) def create_db( self, total_instances: int, db_subnet_group_name: str, node_type: str, security_group_id: str, engine_version: str = None, backup_retention_days: int = 30, public: bool = False, tags: dict = None, ): if not engine_version: engine_version = self.db_engine_version cluster_id = self.cluster_id_format.format(uuid.uuid4().hex) master_creds = Credentials.new() logger.info( f"Creating {self.db_type} cluster", extra=dict( cluster_id=cluster_id, engine_version=engine_version, db_subnet_group_name=db_subnet_group_name, security_group_id=security_group_id, ), ) db_result = self.rds.create_db_cluster( DBClusterIdentifier=cluster_id, Engine=self.db_engine, EngineVersion=engine_version, MasterUsername=master_creds.username, MasterUserPassword=master_creds.password, Tags=[{"Key": k, "Value": v} for k, v in (tags or {}).items()], DBSubnetGroupName=db_subnet_group_name, VpcSecurityGroupIds=[security_group_id], BackupRetentionPeriod=backup_retention_days, )["DBCluster"] endpoint = db_result["Endpoint"] reader_endpoint = db_result["ReaderEndpoint"] port = db_result["Port"] arn = db_result["DBClusterArn"] for i in range(total_instances): instance_id = f"{cluster_id}-Instance-{i}" logger.info( "Creating instances on the cluster", extra=dict( cluster_id=cluster_id, instance_id=instance_id, node_type=node_type, ), ) instance_result = self.rds.create_db_instance( DBInstanceIdentifier=instance_id, DBClusterIdentifier=cluster_id, DBInstanceClass=node_type, Engine=self.db_engine, PubliclyAccessible=public, Tags=[{"Key": k, "Value": v} for k, v in (tags or {}).items()], ) tenant_database_name = self.tenant_database_name_format.format(uuid.uuid4().hex) db = DbInstance( cluster_id=cluster_id, arn=arn, endpoint=endpoint, reader_endpoint=reader_endpoint, port=port, database_name=tenant_database_name, master_user=master_creds, ) print(db) # TODO: remove logger.info("Waiting until DB instances are online. This will take a while (~5 min)") logger.info("Connecting DB instance using master credentials", extra=dict(endpoint=db.endpoint)) self._check_connection( connection=db, credentials=master_creds, ) tenant_creds = Credentials.new() logger.info( "Connection successful. Creating a new database and user", extra=dict( endpoint=endpoint, username=tenant_creds.username, database_name=tenant_database_name, ), ) self.create_tenant( connection=db, master_user=master_creds, database_name=tenant_database_name, tenant_user=tenant_creds, ) db.user = tenant_creds return db @retry(timeout=10 * 60) def _check_connection(self, connection: DbConnection, credentials: Credentials, database_name: str = "postgres"): psycopg2.connect( host=connection.endpoint, port=connection.port, dbname=database_name, user=credentials.username, password=credentials.password, connect_timeout=15, ).close() def _find_instance(self, endpoint: str) -> typing.Optional[DbInstance]: func = functools.partial( self.rds.describe_db_clusters, Filters=[{"Name": "engine", "Values": [self.db_engine]}], MaxRecords=100 ) cluster = None for it in paginate_by_marker(func, "DBClusters"): if it["Endpoint"] == endpoint: cluster = it break if cluster: return DbInstance( endpoint=cluster["Endpoint"], port=cluster["Port"], cluster_id=cluster["DBClusterIdentifier"], arn=cluster["DBClusterArn"], reader_endpoint=cluster["ReaderEndpoint"], ) return None def create_tenant( self, connection: DbConnection, master_user: Credentials, tenant_user: Credentials, database_name: typing.Optional[str] = None, ) -> DbInstance: if not database_name: database_name = self.tenant_database_name_format.format(uuid.uuid4().hex) try: con = psycopg2.connect( host=connection.endpoint, port=connection.port, dbname=self.default_master_database_name, user=master_user.username, password=master_user.password, connect_timeout=30, ) con.set_isolation_level(psycopg2.extensions.ISOLATION_LEVEL_AUTOCOMMIT) cur = con.cursor() cur.execute( sql.SQL("CREATE DATABASE {}").format( sql.Identifier(database_name), ), ) cur.execute( sql.SQL("CREATE USER {} WITH ENCRYPTED PASSWORD {}").format( sql.Identifier(tenant_user.username), sql.Placeholder(), ), [tenant_user.password], ) cur.execute( sql.SQL("GRANT ALL PRIVILEGES ON DATABASE {} TO {}").format( sql.Identifier(database_name), sql.Identifier(tenant_user.username), ) ) # `with` block would create an implicit transaction con.close() db = self._find_instance(connection.endpoint) db.database_name = database_name db.user = tenant_user db.master_user = master_user return db except psycopg2.Error: logging.error( f"Failed to create database and user", exc_info=True, extra=dict(endpoint=connection.endpoint, database_name=database_name), ) raise class MysqlService(PostgresqlService): db_type = "MySQL" db_engine = "aurora-mysql" db_engine_version = "5.7.12" default_master_database_name = "mysql" cluster_id_format = "mysql-{}" # must start with a letter @retry(timeout=10 * 60) def _check_connection(self, connection: DbConnection, credentials: Credentials, database_name: str = "postgres"): pymysql.connect( host=connection.endpoint, port=connection.port, user=credentials.username, password=credentials.password, connect_timeout=15, ).close() def create_tenant( self, connection: DbConnection, master_user: Credentials, tenant_user: Credentials, database_name: typing.Optional[str] = None, ) -> DbInstance: if not database_name: database_name = self.tenant_database_name_format.format(uuid.uuid4().hex) try: with pymysql.connect( host=connection.endpoint, port=connection.port, user=master_user.username, password=master_user.password, connect_timeout=15, ) as con: cur = con.cursor() cur.execute(f"CREATE DATABASE {database_name}") cur.execute(f"CREATE USER {tenant_user.username} IDENTIFIED BY '{tenant_user.password}'") cur.execute( f"GRANT ALL PRIVILEGES ON {database_name}.* TO {tenant_user.username}@'%' IDENTIFIED BY '{tenant_user.password}'" ) db = self._find_instance(endpoint=connection.endpoint) db.database_name = database_name db.master_user = master_user db.user = tenant_user return db except pymysql.Error: logger.error( "Failed to create database and user", exc_info=True, extra=dict(endpoint=connection.endpoint, database_name=database_name), ) raise if __name__ == "__main__": logging.basicConfig(level=logging.INFO, format=f"%(asctime)s: {logging.BASIC_FORMAT}") logging.getLogger("botocore").setLevel(logging.INFO) # m = MysqlService("eu-central-1") # db = m.create_db( # total_instances=1, # db_subnet_group_name="mydbsubnet", # security_group_id="sg-c11a77ac", # node_type="db.t3.medium", # tags={"CostCenter": "hello"}, # public=True, # ) # print(db) # p = PostgresqlService("eu-central-1") # db = p.create_db( # total_instances=1, # db_subnet_group_name="mydbsubnet", # security_group_id="sg-c11a77ac", # node_type="db.t3.medium", # tags={"CostCenter": "hello"}, # public=True, # ) # p.create_tenant( # connection=DbConnection( # endpoint="pg0a085614d5eb4f34affd2a24b1fb18c4.cluster-c8v5rp0ouaey.eu-central-1.rds.amazonaws.com", # port=5432, # ), # master_user=Credentials( # username="u7d0f9621f5a5419288a4ef7ce9fd5fe", password="ca9473879b794dd3a0b0f0d2a8d4aebb" # ), # tenant_user=Credentials.new(), # database_name="mydb", # )