"""Lambda extract-native-configuration function module.""" import json from typing import Any, Dict, List, Optional import boto3 import sentry_sdk from lambdacommon.common_config import logger from sentry_sdk.integrations.aws_lambda import AwsLambdaIntegration import config from common.logic import database # noqa from src.logic import mysql, postgresql if config.SENTRY_DSN: sentry_sdk.init( dsn=config.SENTRY_DSN, environment=config.ENVIRONMENT, integrations=[AwsLambdaIntegration(timeout_warning=True)], ) def assume_source_account_role(account_id: str, role_name: str) -> Dict[str, str]: """Assume IAM role in source account. Args: account_id (str): AWS account ID role_name (str): The role to assume Returns: dict: dict of AWS credentials """ client = boto3.client("sts") assume_role_response = client.assume_role( RoleArn=f"arn:aws:iam::{account_id}:role/{role_name}", RoleSessionName=config.SERVICE_NAME, ExternalId=config.EXTERNAL_ID, DurationSeconds=3600, ) credentials = assume_role_response["Credentials"] return credentials def get_cluster_parameter_group_name( cluster_identifier: str, client: boto3.client ) -> Optional[str]: """Return the parameter group name for the given cluster identifier. Args: cluster_identifier: The identifier of the RDS cluster. client: Boto3 RDS client. Returns: The name of the cluster parameter group, or None if not found. """ response = client.describe_db_clusters(DBClusterIdentifier=cluster_identifier) clusters = response["DBClusters"] if not clusters: return None return clusters[0]["DBClusterParameterGroup"] def get_modified_parameters( engine: str, parameter_group_name: str, client: boto3.client ) -> List[Dict[str, Any]]: """Return a list of modified parameters in the given parameter group. Args: engine: The database engine type. parameter_group_name: The name of the parameter group. client: Boto3 RDS client. Returns: List of modified parameters with their names and values. """ omitted_patterns = config.OMITTED_PARAMETER_GROUP_PATTERNS.get(engine, ()) renamed_params = config.RENAMED_PARAMETER_GROUP_NAMES.get(engine, {}) paginator = client.get_paginator("describe_db_cluster_parameters") modified_parameters = [] for page in paginator.paginate( DBClusterParameterGroupName=parameter_group_name ): for param in page["Parameters"]: if ( param.get("Source") == "user" or param["ParameterName"] in config.REQUIRED_PARAMETER_GROUP_VALUES ): if ( not param["ParameterName"] .lower() .startswith(omitted_patterns) ): modified_parameters.append( { "ParameterName": renamed_params.get( param["ParameterName"], param["ParameterName"] ), "ParameterValue": param.get("ParameterValue"), } ) return modified_parameters def write_modified_parameters_to_file( modified_params: List[Dict[str, Any]], db_engine: str ) -> str: """Write modified parameters to a configuration file. Args: modified_params: List of modified parameters. db_engine: The database engine type ("mysql" or "postgresql"). Returns: Path to the options file. """ header = None if db_engine == "mysql": header = "[mysqld]\n" options_file = "/tmp/my.cnf" elif db_engine == "postgresql": options_file = "/tmp/postgresql.conf" else: raise ValueError(f"Unsupported db_engine: {db_engine}") with open(options_file, "w") as f: if header and modified_params: f.write(header) for param in sorted(modified_params, key=lambda x: x["ParameterName"]): f.write(f"{param['ParameterName']}={param['ParameterValue']}\n") return options_file def handler(event: Dict[str, Any], context: Any) -> Dict[str, str]: """Lambda entry point. Args: event: Lambda event payload containing db_name and db_type. context: Lambda runtime context object. Returns: Dict containing completion status. """ s3_client = boto3.client("s3") try: db_name = event.get("db_name") db_type = event.get("db_type") source_account_id = event.get("source_account_id") if not db_name or not db_type or not source_account_id: logger.error( "Missing 'db_name', 'db_type', or 'source_account_id' in payload" ) raise ValueError( "Missing 'db_name', 'db_type', or 'source_account_id' in payload" ) source_account_role = event.get( "source_account_role", config.CROSS_ACCOUNT_BACKUP_ROLE_NAME ) logger.info( f"Assuming role {source_account_role} in account {source_account_id}" ) credentials = assume_source_account_role( source_account_id, source_account_role ) rds_client = boto3.client( "rds", region_name=config.AWS_DEFAULT_REGION, aws_access_key_id=credentials["AccessKeyId"], aws_secret_access_key=credentials["SecretAccessKey"], aws_session_token=credentials["SessionToken"], ) logger.info( f"Resetting master credentials for {db_name}, type: {db_type}" ) db_credentials = database.reset_rds_master_credentials( db_name, db_type, wait=True, client=rds_client ) db_config = database.get_config(db_name, db_type, client=rds_client) db_name_s3_prefix = db_name.partition("-")[2] if "mysql" in db_config["Engine"]: schema_file = mysql.get_database_ddl(db_credentials) logger.info(f"Schema file generated: {schema_file}") db_engine = "mysql" elif "postgresql" in db_config["Engine"]: schema_file = postgresql.get_database_ddl(db_credentials) logger.info(f"Schema file generated: {schema_file}") db_engine = "postgresql" else: raise ValueError( f"Unsupported database engine: {db_config['Engine']}" ) s3_client.upload_file( schema_file, config.S3_BUCKET_NAME, f"artifacts/{db_name_s3_prefix}/schema.sql", ) engine_version = ".".join( db_config.get("EngineVersion", "").split(".")[0:2] ) engine_object = json.dumps( { "engine": db_engine, "version": engine_version, } ) s3_client.put_object( Bucket=config.S3_BUCKET_NAME, Key=f"artifacts/{db_name_s3_prefix}/engine.json", Body=engine_object.encode("utf-8"), ) parameter_group = get_cluster_parameter_group_name(db_name, rds_client) if parameter_group: modified_parameters = get_modified_parameters( db_engine, parameter_group, rds_client ) options_file = write_modified_parameters_to_file( modified_parameters, db_engine ) logger.info(f"Options file generated: {options_file}") options_file_name = options_file.split("/")[-1] s3_client.upload_file( options_file, config.S3_BUCKET_NAME, f"artifacts/{db_name_s3_prefix}/{options_file_name}", ) except Exception as e: logger.exception(str(e)) raise e return {"status": "completed"}