Files
2024-12-24 20:49:41 +03:00

402 lines
13 KiB
Python

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",
# )