"""Entrypoint.""" import argparse import os import pymysql user = os.environ['DB_USER'] password = os.environ['DB_PASSWORD'] host = os.environ['DB_HOST'] schema = os.environ['DB_SCHEMA'] port = os.environ.get('DB_PORT', 3306) DB_CONN = pymysql.connect( host=host, user=user, password=password, database=schema, cursorclass=pymysql.cursors.DictCursor ) def main(): """Start.""" parser = argparse.ArgumentParser( formatter_class=argparse.ArgumentDefaultsHelpFormatter ) parser.add_argument( 'table', type=str, help='name of table to anchor on' ) parser.add_argument( '--column', type=str, help='column to search for data on' ) parser.add_argument( '--values', type=str, default=None, nargs='+', help='values to filter column on' ) args = parser.parse_args() gather( args.table, args.column, args.values ) def gather(table, column=None, values=None): """Recursively generate SQL statement.""" rows = query_rows(table, column, values) if values else [] # gather CREATE TABLE and INSERT for dependant data for mapping in query_fk_mappings(table): # fk rows attached to rows fk_ids = list(set([ x[mapping['base_column']] for x in rows ])) # recurvisely search through fk table gather( mapping['fk_table'], mapping['fk_column'], [x for x in fk_ids if x is not None] ) # output CREATE TABLE statement print(query_create_table(table)) # output INSERT statements, if rows needed if rows: print(rows_to_insert(table, rows)) def _format_value_for_insert(value): """Change value to be value for INSERT.""" if value is None: return 'null' value = str(value) value = value.replace("'", "\\'") value = "'" + value + "'" return value def rows_to_insert(table, rows): """Convert row dict to INSERT statement.""" columns = rows[0].keys() values = [ [ _format_value_for_insert(row[column_name]) for column_name in columns ] for row in rows ] column_stmt = '(' + ','.join(columns) + ')' values_stmt = ',\n'.join( [ '(' + ','.join(value) + ')' for value in values ] ) return f""" INSERT IGNORE INTO {table} {column_stmt} VALUES {values_stmt} ; """ def query_rows(table, column, values): """Query single table for all rows.""" return query( f""" SELECT * FROM `{_valid(table)}` WHERE `{_valid(column)}` IN %(values)s """, { 'values': values } ) def query_create_table(table): """Query for CREATE TABLE syntax.""" result = query(f'SHOW CREATE TABLE `{_valid(table)}`') stmt = result[0]['Create Table'] + '\n;' return stmt.replace('CREATE TABLE', 'CREATE TABLE IF NOT EXISTS') def query_fk_mappings(table): """Query FK references for a table.""" return query( """ SELECT TABLE_NAME AS base_table, COLUMN_NAME AS base_column, REFERENCED_TABLE_NAME AS fk_table, REFERENCED_COLUMN_NAME AS fk_column FROM INFORMATION_SCHEMA.KEY_COLUMN_USAGE WHERE REFERENCED_TABLE_SCHEMA = (SELECT DATABASE()) AND TABLE_NAME = %(table)s """, { 'table': table } ) def _valid(string): if not string.replace('_', '').isalnum(): raise Exception(f'{string} not valid for direct query injection') return string def query(sql, params=dict()): """Perform query.""" with DB_CONN.cursor() as cursor: cursor.execute(sql, params) results = cursor.fetchall() return results if __name__ == '__main__': main()