import logging import os from dotenv import load_dotenv from fabric import Connection from jinja2 import Template load_dotenv() logger = logging.getLogger(__name__) logging.basicConfig( level=logging.INFO, format='%(asctime)s %(levelname)s [%(name)s]: %(message)s' ) BATCH_ID = os.environ.get('BATCH_ID') SSH_HOST = os.environ.get('SSH_HOST') SSH_USER = os.environ.get('SSH_USER') SSH_KEY_FILENAME = os.environ.get('SSH_KEY_FILENAME') OUTFILE_DIR = os.environ.get('OUTFILE_DIR') SERVICE_USER = os.environ.get('SERVICE_USER') SERVICE_GROUP = os.environ.get('SERVICE_GROUP') MYSQL_USER = os.environ.get('MYSQL_USER') MYSQL_PASSWORD = os.environ.get('MYSQL_PASSWORD') SPLIT_SIZE = os.environ.get('SPLIT_SIZE', '10G') S3_BUCKET = os.environ.get('S3_BUCKET') S3_KEY_ID = os.environ.get('S3_KEY_ID') S3_SECRET_KEY = os.environ.get('S3_SECRET_KEY') EXTRACT_LIMIT = int(os.environ.get('EXTRACT_LIMIT', 10)) QUERY_TEMPLATE = Template(""" SELECT * FROM sales_file_delivery.dig_sales_testfile_abacus_test WHERE batch_id = '{{batch_id}}' LIMIT {{limit}} """) def main(): try: logger.info(f'Starting extract_sales for batch {BATCH_ID}') raw_dir = f'{OUTFILE_DIR}/distro/{BATCH_ID}/raw' parts_dir = f'{OUTFILE_DIR}/distro/{BATCH_ID}/parts' outfile_path = f'{raw_dir}/outfile.txt' parts_path = f'{parts_dir}/outfile_parts_' s3_path = f's3://{S3_BUCKET}/stmtdb-to-sf/extract-sales/distro/{BATCH_ID}/' query = QUERY_TEMPLATE.render(batch_id=BATCH_ID, limit=EXTRACT_LIMIT) ssh_conn = Connection( host=SSH_HOST, user=SSH_USER, connect_kwargs={ 'key_filename': SSH_KEY_FILENAME } ) # Create `raw` dir ssh_conn.run( f'install -dv -m 0775 -o {SERVICE_USER} -g {SERVICE_GROUP} {raw_dir}', hide=True ) # Create `parts` dir ssh_conn.run( f'install -dv -m 0775 -o {SERVICE_USER} -g {SERVICE_GROUP} {parts_dir}', hide=True ) logger.info('Created directories') # Select into outfile # NOTE: Using `-N` to not include column names in the outfile for easier # loading into Snowflake # NOTE: Using `--quick` to print each row as it is received to not overload the server # by keeping the whole result set in memory ssh_conn.run( f'mysql --quick -N -u{MYSQL_USER} -p{MYSQL_PASSWORD} -e "{query}" > {outfile_path}', hide=True ) logger.info('Generated outfile') # # Get number of lines in outfile # # NOTE: Using `ripgrep` for better performance # line_result = ssh_conn.run(f'rg -c ^ {outfile_path}', hide=True) # line_count = int(line_result.stdout.strip()) # # logger.info(f'The outfile contains {line_count} lines') # # split_size = SPLIT_SIZE if line_count > SPLIT_SIZE else line_count # Split and compress # NOTE: Using `pigz` for better performance ssh_conn.run( f'split -C {SPLIT_SIZE} --filter=\'/usr/bin/pigz > $FILE.gz\' {outfile_path} {parts_path}', hide=True ) logger.info('Split and compressed outfile') # Upload to S3 # NOTE: Using `s5cmd` for better performance aws_creds = f'AWS_ACCESS_KEY_ID={S3_KEY_ID}' aws_creds += f' AWS_SECRET_ACCESS_KEY={S3_SECRET_KEY}' ssh_conn.run(f'{aws_creds} s5cmd cp {parts_dir} {s3_path}', hide=True) logger.info(f'Parts uploaded to {s3_path}') # Clean up ssh_conn.run(f'rm -rf {raw_dir} {parts_dir}', hide=True) logger.info(f'Files deleted') except Exception as e: logger.error(f'Error in extract_sales: {e}') raise e if __name__ == '__main__': main()