From 81d8494755e453d9ce6b39ff259ddb30100f7bce Mon Sep 17 00:00:00 2001 From: Ken Tsang Date: Wed, 2 Aug 2017 10:17:21 +0100 Subject: [PATCH] Add tests for s3_client --- app/main/s3_client.py | 96 +++++++++++++++++--------------- tests/app/main/test_s3_client.py | 95 +++++++++++++++++++++++++++++++ 2 files changed, 145 insertions(+), 46 deletions(-) create mode 100644 tests/app/main/test_s3_client.py diff --git a/app/main/s3_client.py b/app/main/s3_client.py index 64efa11c8..75404ac5e 100644 --- a/app/main/s3_client.py +++ b/app/main/s3_client.py @@ -1,6 +1,6 @@ import uuid import botocore -from boto3 import resource, client +from boto3 import resource from flask import current_app from notifications_utils.s3 import s3upload as utils_s3upload @@ -9,6 +9,33 @@ TEMP_TAG = 'temp-{user_id}_' LOGO_LOCATION_STRUCTURE = '{temp}{unique_id}-{filename}' +def get_s3_object(bucket_name, filename): + s3 = resource('s3') + return s3.Object(bucket_name, filename) + + +def delete_s3_object(filename): + bucket_name = current_app.config['LOGO_UPLOAD_BUCKET_NAME'] + get_s3_object(bucket_name, filename).delete() + + +def rename_s3_object(old_name, new_name): + bucket_name = current_app.config['LOGO_UPLOAD_BUCKET_NAME'] + get_s3_object(bucket_name, new_name).copy_from( + CopySource='{}/{}'.format(bucket_name, old_name)) + delete_s3_object(old_name) + + +def get_s3_objects_filter_by_prefix(prefix): + bucket_name = current_app.config['LOGO_UPLOAD_BUCKET_NAME'] + s3 = resource('s3') + return s3.Bucket(bucket_name).objects.filter(Prefix=prefix) + + +def get_temp_truncated_filename(filename, user_id): + return filename[len(TEMP_TAG.format(user_id=user_id)):] + + def s3upload(service_id, filedata, region): upload_id = str(uuid.uuid4()) upload_file_name = FILE_LOCATION_STRUCTURE.format(service_id, upload_id) @@ -22,10 +49,9 @@ def s3upload(service_id, filedata, region): def s3download(service_id, upload_id): contents = '' try: - s3 = resource('s3') bucket_name = current_app.config['CSV_UPLOAD_BUCKET_NAME'] upload_file_name = FILE_LOCATION_STRUCTURE.format(service_id, upload_id) - key = s3.Object(bucket_name, upload_file_name) + key = get_s3_object(bucket_name, upload_file_name) contents = key.get()['Body'].read().decode('utf-8') except botocore.exceptions.ClientError as e: current_app.logger.error("Unable to download s3 file {}".format( @@ -40,59 +66,37 @@ def upload_logo(filename, filedata, region, user_id): unique_id=str(uuid.uuid4()), filename=filename ) - utils_s3upload(filedata=filedata, - region=region, - bucket_name=current_app.config['LOGO_UPLOAD_BUCKET_NAME'], - file_location=upload_file_name, - content_type='image/png') + bucket_name = current_app.config['LOGO_UPLOAD_BUCKET_NAME'] + utils_s3upload( + filedata=filedata, + region=region, + bucket_name=bucket_name, + file_location=upload_file_name, + content_type='image/png' + ) + return upload_file_name def persist_logo(filename, user_id): - try: - if filename.startswith(TEMP_TAG.format(user_id=user_id)): - persisted_filename = filename[len(TEMP_TAG.format(user_id)):] - else: - return filename + if filename.startswith(TEMP_TAG.format(user_id=user_id)): + persisted_filename = get_temp_truncated_filename( + filename=filename, user_id=user_id) + else: + return filename - s3 = resource('s3') - bucket_name = current_app.config['LOGO_UPLOAD_BUCKET_NAME'] + rename_s3_object(filename, persisted_filename) - s3.Object(bucket_name, persisted_filename).copy_from(CopySource='{}/{}'.format(bucket_name, filename)) - s3.Object(bucket_name, filename).delete() - - return persisted_filename - except botocore.exceptions.ClientError as e: - current_app.logger.error("Unable to get s3 bucket contents {}".format( - bucket_name)) - raise e + return persisted_filename def delete_temp_files_created_by(user_id): - try: - s3 = resource('s3') - bucket_name = current_app.config['LOGO_UPLOAD_BUCKET_NAME'] - - for obj in s3.Bucket(bucket_name).objects.filter(Prefix=TEMP_TAG.format(user_id)): - s3.Object(bucket_name, obj.key).delete() - - except botocore.exceptions.ClientError as e: - current_app.logger.error("Unable to delete s3 bucket temp files created by {} from {}".format( - user_id, bucket_name)) - raise e + for obj in get_s3_objects_filter_by_prefix(TEMP_TAG.format(user_id=user_id)): + delete_s3_object(obj.key) def delete_temp_file(filename): - try: - if not filename.startswith(TEMP_TAG): - raise ValueError('Not a temp file') + if not filename.startswith(TEMP_TAG[:5]): + raise ValueError('Not a temp file: {}'.format(filename)) - s3 = resource('s3') - bucket_name = current_app.config['LOGO_UPLOAD_BUCKET_NAME'] - - s3.Object(bucket_name, filename).delete() - - except botocore.exceptions.ClientError as e: - current_app.logger.error("Unable to delete s3 bucket file {} from {}".format( - filename, bucket_name)) - raise e + delete_s3_object(filename) diff --git a/tests/app/main/test_s3_client.py b/tests/app/main/test_s3_client.py new file mode 100644 index 000000000..d86e03b20 --- /dev/null +++ b/tests/app/main/test_s3_client.py @@ -0,0 +1,95 @@ +from collections import namedtuple +from unittest.mock import call +import pytest + +from app.main.s3_client import ( + upload_logo, + persist_logo, + delete_temp_file, + delete_temp_files_created_by, + get_temp_truncated_filename, + LOGO_LOCATION_STRUCTURE, + TEMP_TAG +) + +bucket = 'test_bucket' +data = {'data': 'some_data'} +filename = 'test.png' +upload_id = 'test_uuid' +region = 'eu-west1' + + +@pytest.fixture +def upload_filename(fake_uuid): + return LOGO_LOCATION_STRUCTURE.format( + temp=TEMP_TAG.format(user_id=fake_uuid), unique_id=upload_id, filename=filename) + + +def test_upload_logo_calls_correct_args(client, mocker, fake_uuid, upload_filename): + mocker.patch('uuid.uuid4', return_value=upload_id) + mocker.patch.dict('flask.current_app.config', {'LOGO_UPLOAD_BUCKET_NAME': bucket}) + mocked_s3_upload = mocker.patch('app.main.s3_client.utils_s3upload') + + upload_logo(filename=filename, user_id=fake_uuid, filedata=data, region=region) + + assert mocked_s3_upload.called_once_with( + filedata=data, + region=region, + file_location=upload_filename, + bucket_name=bucket + ) + + +def test_persist_logo(client, mocker, fake_uuid, upload_filename): + mocker.patch.dict('flask.current_app.config', {'LOGO_UPLOAD_BUCKET_NAME': bucket}) + mocked_rename_s3_object = mocker.patch('app.main.s3_client.rename_s3_object') + + persisted_filename = persist_logo(filename=upload_filename, user_id=fake_uuid) + + assert mocked_rename_s3_object.called_once_with( + upload_filename, get_temp_truncated_filename(upload_filename, fake_uuid)) + assert persisted_filename == get_temp_truncated_filename(upload_filename, fake_uuid) + + +def test_persist_logo_returns_if_not_temp(client, mocker, fake_uuid): + filename = 'logo.png' + mocker.patch.dict('flask.current_app.config', {'LOGO_UPLOAD_BUCKET_NAME': bucket}) + mocked_rename_s3_object = mocker.patch('app.main.s3_client.rename_s3_object') + + persisted_filename = persist_logo(filename=filename, user_id=fake_uuid) + + assert not mocked_rename_s3_object.called + assert persisted_filename == filename + + +def test_delete_temp_files_created_by_user(client, mocker, fake_uuid): + obj = namedtuple("obj", ["key"]) + objs = [obj(key='test1'), obj(key='test2')] + + mocker.patch('app.main.s3_client.get_s3_objects_filter_by_prefix', return_value=objs) + mocked_delete_s3_object = mocker.patch('app.main.s3_client.delete_s3_object') + + delete_temp_files_created_by(fake_uuid) + + assert mocked_delete_s3_object.called_with_args(objs[0].key) + for index, arg in enumerate(mocked_delete_s3_object.call_args_list): + assert arg == call(objs[index].key) + + +def test_delete_single_temp_file(client, mocker, fake_uuid, upload_filename): + mocked_delete_s3_object = mocker.patch('app.main.s3_client.delete_s3_object') + + delete_temp_file(upload_filename) + + assert mocked_delete_s3_object.called_with_args(upload_filename) + + +def test_does_not_delete_non_temp_file(client, mocker, fake_uuid): + filename = 'logo.png' + mocked_delete_s3_object = mocker.patch('app.main.s3_client.delete_s3_object') + + with pytest.raises(ValueError) as error: + delete_temp_file(filename) + + assert mocked_delete_s3_object.called_with_args(filename) + assert str(error.value) == 'Not a temp file: {}'.format(filename)