import time import pymysql from accounting_run_utils import config logger = config.logger def get_mysql_connection() -> pymysql.connections.Connection: connection = pymysql.connect( host=config.MYSQL_HOST, user=config.MYSQL_USER, password=config.MYSQL_PASSWORD, database=config.MYSQL_DATABASE, charset="utf8", cursorclass=pymysql.cursors.DictCursor, ) return connection def get_table_partition_count(database: str, table: str) -> list: connection = get_mysql_connection() logger.debug("Got a MySQL connection") with connection.cursor() as cursor: sql = f"SELECT PARTITION_NAME from information_schema.PARTITIONS where TABLE_SCHEMA = '{ database}' and TABLE_NAME = '{table}';" cursor.execute(sql) result = cursor.fetchall() partition_count = [partition["PARTITION_NAME"] for partition in result] logger.info(f"Partition list: {partition_count}") if connection.show_warnings(): logger.warning(f"Warnings: {connection.show_warnings()}") connection.close() return partition_count def get_table_subpartition_count_for_period_id( database: str, table: str, period_id: str ) -> list: """ Create table DDL uses VALUES LESS THAN (period_id) for partitioning, so use (period_id + 1) as the lookup value. """ partition_description = int(period_id) + 1 connection = get_mysql_connection() logger.debug("Got a MySQL connection") with connection.cursor() as cursor: sql = f"SELECT SUBPARTITION_NAME from information_schema.PARTITIONS where TABLE_SCHEMA = '{ database}' and TABLE_NAME = '{table}' and PARTITION_DESCRIPTION = '{ partition_description}' order by TABLE_ROWS desc;" cursor.execute(sql) result = cursor.fetchall() subpartition_count = [ partition["SUBPARTITION_NAME"] for partition in result ] logger.info(f"Subpartition list: {subpartition_count}") if connection.show_warnings(): logger.warning(f"Warnings: {connection.show_warnings()}") connection.close() return subpartition_count def load_data(sql: str, object_key: str) -> int: connection = get_mysql_connection() logger.debug("Got a MySQL connection") with connection.cursor() as cursor: logger.info(f"Loading data from {object_key}") now = time.time() result = cursor.execute(sql) post_query = time.time() logger.info(f"Updated rows from {object_key}: {result}") logger.debug(f"Load time for {object_key} is {post_query - now}") connection.commit() if connection.show_warnings(): logger.warning(f"Warnings: {connection.show_warnings()}") connection.close() return result def prepare_partitioned_query( absolute_query_file_path: str, partition: str ) -> str: with open(absolute_query_file_path) as sql_file: prepared_statement = sql_file.read() """ Partition names in queries cannot contain quotes, so replace the partition name in the prepared statement with the actual partition to work around that limitation. """ if "PARTITION_NAME" in prepared_statement: sql = prepared_statement.replace("PARTITION_NAME", partition) else: raise ValueError( 'SQL file must contain the phrase "PARTITION_NAME" to be replaced' ) return sql def run_query_with_no_result(sql: str) -> None: """Trivial wrapper function to log query time.""" connection = get_mysql_connection() logger.debug("Got a MySQL connection") with connection.cursor() as cursor: logger.info(f"Running query {sql}") now = time.time() cursor.execute(sql) post_query = time.time() logger.info(f"Query time for {sql} is {post_query - now}") connection.commit() if connection.show_warnings(): logger.warning(f"Warnings: {connection.show_warnings()}") connection.close() def select_data_into_outfile( sql: str, object_key: str, sql_args: dict | None = None, ) -> int: print(sql_args) connection = get_mysql_connection() logger.debug("Got a MySQL connection") with connection.cursor() as cursor: logger.info(f"Selecting data into outfile {object_key}") now = time.time() result = cursor.execute(sql, args=sql_args) post_query = time.time() logger.info(f"Updated rows into outfile {object_key}: {result}") logger.info(f"Query time for {sql} is {post_query - now}\n") connection.commit() if connection.show_warnings(): logger.warning(f"Warnings: {connection.show_warnings()}") connection.close() return result