import csv import functools import io from unittest.mock import Mock from unittest.mock import patch from moto import mock_aws import boto3 from flexmock import flexmock from oto import response import pytest from masters_registry import config from masters_registry.constant import bulk_tasks_const from masters_registry.constant import error from masters_registry.constant import field_const from masters_registry.logic import bulk_tasks as logic from masters_registry.models import bulk_tasks from masters_registry.connectors.s3 import client as outside_s3_client from tests import test_utils from tests.helpers import patches from tests.logic.bulk_logic_fixtures import task_info_lock_report from tests.logic.bulk_logic_fixtures import task_info_update_report from tests.logic.bulk_logic_fixtures import task_initial_db from tests.logic.bulk_logic_fixtures import task_update_db # noqa from tests.logic.bulk_logic_fixtures import upc_json_for_done_status from tests.logic.bulk_logic_fixtures import upc_json_for_failed_status from tests.logic.bulk_logic_fixtures import upc_json_for_partial_fail_status s3_resource = boto3.resource("s3") def test_bulk_tasks(monkeypatch): """Test bulk_tasks function """ num_records = 10 order_by = 'order' order_direction = 'direction' get_task_statuses_count = Mock(return_value=3) monkeypatch.setattr( bulk_tasks, 'get_task_statuses_count', get_task_statuses_count) get_task_report_result = [1, 2, 3] get_task_report = Mock(return_value=get_task_report_result) monkeypatch.setattr( bulk_tasks, 'get_task_report', get_task_report) response = logic.get_tasks(num_records, order_by, order_direction, 50) get_task_statuses_count.assert_called_once_with() get_task_report.assert_called_with( num_records, order_by, order_direction, 50) assert response.message == { field_const.ITEMS: get_task_report_result, field_const.TOTAL_COUNT: 3 } def test_generate_report_failed_invalid_correlation_id( bulk_status_db_fixture, valid_correlation_id): invalid_correlation_id = 'invalid_correlation_id' response = logic.generate_report( invalid_correlation_id, valid_correlation_id) assert response.errors['message'] == error.BULK_TASK_DOES_NOT_EXIST @mock_aws def test_generate_report_success_status(task_update_db, valid_correlation_id, feature_engine): # noqa """Testing generate_report function. Successful case. Checking 200 status of Response. """ from moto.core import patch_client, patch_resource patch_client(outside_s3_client) patch_resource(s3_resource) outside_s3_client.create_bucket(Bucket=config.REPORTS_BUCKET_NAME) (flexmock(bulk_tasks) .should_receive('get_task') .and_return(task_update_db)) resp = logic.generate_report(2, valid_correlation_id) assert resp.status == 200 def test_generate_report_upload_to_s3(task_update_db, valid_correlation_id): # noqa """Testing generate_report function far a case when upload_to_s3 is True """ (flexmock(bulk_tasks) .should_receive('get_task') .and_return(task_update_db)) expected_url = 'https://dev-orcdbucket.s3.amazonaws.com/file.csv' expected_file_name = '2000-12-12_Bulk_Update_UPCs_UPC_12.csv' import_report_type = 'upc' (flexmock(logic) .should_receive('upload_report') .with_args(expected_file_name, functools.partial) .and_return(expected_url)) resp = logic.generate_report( task_id=2, correlation_id=valid_correlation_id, import_report_type=import_report_type, ) assert resp == expected_url def test_generate_report_task_not_found( bulk_status_db_fixture, valid_correlation_id, get_task_not_found): (flexmock(bulk_tasks) .should_receive('get_task') .and_return(get_task_not_found)) resp = logic.generate_report(777, valid_correlation_id) assert resp.status == 404 @pytest.mark.parametrize('upc_json, report_type, expectation', [ ( upc_json_for_done_status(), 'upc', bulk_tasks_const.UPDATE_STATUS_SUCCESS ), ( upc_json_for_partial_fail_status(), 'upc', bulk_tasks_const.UPDATE_STATUS_PARTIAL_FAIL ), ( upc_json_for_partial_fail_status(), 'isrc', bulk_tasks_const.UPDATE_STATUS_FAIL ), ( upc_json_for_failed_status(), 'upc', bulk_tasks_const.UPDATE_STATUS_FAIL ) ]) def test_calculate_status(upc_json, report_type, expectation): """Test _calculate_status function from bulk tasks logic. """ (flexmock(logic) .should_receive('_is_claimed_by_another_owner') .and_return(False)) result_status = logic._calculate_status(upc_json, report_type) assert result_status == expectation @pytest.mark.parametrize('task_info, import_type, expectation', [ ( task_info_update_report(), 'isrc', '2000-12-12_Bulk_Update_UPCs_ISRC_15.csv' ), ( task_info_update_report(), 'upc', '2000-12-12_Bulk_Update_UPCs_UPC_15.csv' ), ( task_info_lock_report(), None, '2000-12-12_Bulk_Lock_Territories_ISRC_16.csv' ) ]) def test_create_file_name(task_info, import_type, expectation): from masters_registry.logic.bulk_tasks import _create_file_name result = _create_file_name(task_info, import_type) assert result == expectation def test_generate_initial_report_upload_to_s3(valid_correlation_id): """Testing generate_initial_report function far a case when upload_to_s3 is True """ task = task_initial_db() (flexmock(bulk_tasks) .should_receive('get_task') .and_return(task)) expected_url = 'https://dev-orcdbucket.s3.amazonaws.com/file.csv' expected_file_name = '2000-12-12_Bulk_Update_UPCs_ISRC_123_initial.csv' import_report_type = 'isrc' (flexmock(logic) .should_receive('upload_report') .with_args(expected_file_name, functools.partial) .and_return(expected_url)) resp = logic.generate_initial_report( task_id=2, import_report_type=import_report_type, correlation_id=valid_correlation_id, ) assert resp == expected_url def test_generate_initial_report_failed(valid_correlation_id): """Testing generate_initial_report function. Fail case. Checking for empty response""" failed_get_task_response = response.Response( status=404 ) (flexmock(bulk_tasks) .should_receive('get_task') .and_return(failed_get_task_response)) resp = logic.generate_initial_report( task_id=666, import_report_type='isrc', correlation_id=valid_correlation_id ) assert resp.status == 404 @patch( 'masters_registry.models.ownership.get_tracks', new=patches.ownership_get_tracks_only_vendor_patch) def test_generate_update_by_isrc_report_internal_conflict(): """Test status set to 'WARNING' if conflit was created""" report = io.StringIO() csv_writer = csv.writer(report) task_info = bulk_tasks.TaskStatus(user_id='123') correlation_id = '12345' upc = 'test_upc' isrc = 'test_isrc' vendor_id = '12345' tuid = '123' error_template = bulk_tasks_const.INTERNAL_CONFLICT_CREATED_WARNING expected_error = '{0}.{1}'.format( error_template.format(countries='PT'), error_template.format(countries='QA')) task_info.result = [ { field_const.ERROR_COUNT: 0, field_const.SUCCESS_COUNT: 1, field_const.UPC: upc, field_const.STATUS_REPORT: [ { field_const.VENDOR_ID: vendor_id, field_const.TUID: tuid, field_const.ISRC: isrc, field_const.SUCCESS: 'True', field_const.ERROR_MESSAGE: '', field_const.CLAIMED_BY_ANOTHER_OWNER: [ ('PT', [11684704]), ('QA', [11432391]) ] } ] } ] expected_row = [ vendor_id, '="{}"'.format(upc), bulk_tasks_const.SUCCESS, bulk_tasks_const.NA, tuid, isrc, bulk_tasks_const.WARNING, expected_error ] logic._generate_update_by_isrc_report( csv_writer, task_info, correlation_id) result_row = list(csv.reader(report.getvalue().split('\n')))[1] assert result_row == expected_row def test_generate_update_by_isrc_report_conflict_resolved(): """Test status set to 'WARNING' if an internal conflict was resolved""" report = io.StringIO() csv_writer = csv.writer(report) task_info = bulk_tasks.TaskStatus(user_id='123') correlation_id = '12345' upc = 'test_upc' isrc = 'test_isrc' tuid = '12345' vendor_id = '123' expected_error = ( bulk_tasks_const.INTERNAL_CONFLICT_RESOLVED_WARNING.format( 'CA,QA')) task_info.result = [ { field_const.ERROR_COUNT: 0, field_const.SUCCESS_COUNT: 1, field_const.UPC: upc, field_const.STATUS_REPORT: [ { field_const.VENDOR_ID: vendor_id, field_const.ISRC: isrc, field_const.TUID: tuid, field_const.SUCCESS: 'True', field_const.ERROR_MESSAGE: '', field_const.RESOLVED_CONFLICT: True, field_const.REMOVED_TERRITORIES: ['CA', 'QA'], field_const.CLAIMED_BY_ANOTHER_OWNER: [] } ] } ] expected_row = [ vendor_id, '="{}"'.format(upc), bulk_tasks_const.SUCCESS, bulk_tasks_const.NA, tuid, isrc, bulk_tasks_const.WARNING, expected_error ] logic._generate_update_by_isrc_report( csv_writer, task_info, correlation_id) result_row = list(csv.reader(report.getvalue().split('\n')))[1] assert result_row == expected_row def test_get_failed_reason_when_conflict_was_resolved(): """Test reason set correctly when a conflict was resolved.""" upc_status = 'Warning' isrc_info = { field_const.ISRC: 'USA370300502', field_const.CLAIMED_BY_ANOTHER_OWNER: [], field_const.RESOLVED_CONFLICT: True, field_const.REMOVED_TERRITORIES: ['CA', 'QA'], field_const.ERROR_MESSAGE: '', field_const.LOCKED_TERRITORIES: ['IT', 'PT'], } result = logic._get_isrc_update_message(isrc_info, upc_status) expected_result = '{0} | {1}'.format( 'Locked Territories: IT,PT', 'Internal conflict resolved for these territories: CA,QA') assert result == expected_result def test_get_isrc_update_message_when_substore_carveouts_has_no_territories(): """Test _get_isrc_update_message when carveouts has no territories""" upc_status = 'Warning' isrc_info = { field_const.ISRC: 'USA370300502', field_const.CLAIMED_BY_ANOTHER_OWNER: [], field_const.RESOLVED_CONFLICT: False, field_const.ERROR_MESSAGE: '', field_const.SUBSTORE_CARVEOUTS: [(453, set())] } result = logic._get_isrc_update_message(isrc_info, upc_status) assert result == '' def test_get_failed_reason_when_claimed_by_another_owner(): """Test reason set correctly when territory claimed by another owner.""" expected_fail_reason = bulk_tasks_const.INTERNAL_CONFLICT_CREATED_WARNING upc_status = 'Warning' tracks = response.Response({182: { # This is only part of dictionary returned by get_track 'isrc': 'USA370300501', 'tuid': 182, 'vendor_id': 77, }}) (flexmock(logic) .should_receive('ownership.get_tracks') .with_args([182]) .and_return(tracks)) error = 'Some error' isrc_info = { field_const.ISRC: 'USA370300502', field_const.CLAIMED_BY_ANOTHER_OWNER: [ ('US', 182), ('UK', 182)], field_const.ERROR_MESSAGE: error } countries = 'US,UK' tuid = '77' result = logic._get_isrc_update_message(isrc_info, upc_status) expected_result = '{error} | {isrc_update_message}'.format( error=error, isrc_update_message=expected_fail_reason.format( tuid=tuid, countries=countries)) assert result == expected_result def test_group_territories_by_tuid(): """Test (country_code, tuid) tuples a grouped correctly.""" data = [('CA', [180]), ('US', 182), ('UK', 182), ('UA', [180, 184])] expected = { 180: ['CA', 'UA'], 182: ['US', 'UK'], 184: ['UA'] } result = logic._group_territories_by_tuid(data) assert result == expected @pytest.fixture def upc_row(): """Row of data from TaskStatus.result column of this services.""" data = { 'error_count': 1, 'success_count': 0, 'upc': 'some upc', 'status_report': [ {'claimed_by_another_owner': [('US', 'some tuid')], 'upc': 'test'} ] } return data def test_calculate_status_sets_partial_fail(upc_row): """ Test status is set to 'Partial fail' if there are territories claimed by another owner. """ status = logic._calculate_status(upc_row, 'upc') assert status == bulk_tasks_const.UPDATE_STATUS_PARTIAL_FAIL @pytest.mark.parametrize( 'claimed_territories, expected_value', [ ([('US', 'some tuid')], True), ([], False) # means there are no territories claimed by someone ]) def test_is_claimed_by_another_owner( claimed_territories, expected_value, upc_row): """ Test isrcs that have conflict because their territories were claimed by another owner can be determined. """ upc_row['status_report'][0]['claimed_by_another_owner'] = ( claimed_territories) claimed_by_another_owner = logic._is_claimed_by_another_owner(upc_row) assert claimed_by_another_owner is expected_value @mock_aws def test_upload_report(): """Test test_upload_report function""" from moto.core import patch_client, patch_resource patch_client(outside_s3_client) patch_resource(s3_resource) outside_s3_client.create_bucket(Bucket=config.REPORTS_BUCKET_NAME) file_name = 'file.csv' report_data = 'abcde' generate_report = Mock(return_value=report_data) key = '{prefix}/{filename}'.format( filename=file_name, prefix=config.REPORTS_FILE_KEY_PREFIX) res = logic.upload_report(file_name, generate_report) report_file = io.BytesIO() outside_s3_client.download_fileobj(Key=key, Fileobj=report_file, Bucket=config.REPORTS_BUCKET_NAME) generate_report.assert_called_with() assert report_file.getvalue().decode() == report_data assert res.message.startswith( 'https://{}.s3.amazonaws.com/{}'.format(config.REPORTS_BUCKET_NAME, key) ) @patch.object(logic.s3, 'check_if_file_exists', new=Mock(return_value=True)) @patch.object(logic.s3, 'get_file_url', new=Mock(return_value='existing_url')) @patch.object(logic.s3, 'upload_file_object', new=Mock()) @patch.object(logic, 'generate_report', new=Mock()) def test_upload_report_file_exists(): """Should return existing file without generating new one.""" upload_result = logic.upload_report( file_name='test_file', generate_report=None) assert upload_result == 'existing_url' expected_file_key = 'masters-registry-reports/bulk-action/test_file' logic.s3.check_if_file_exists.assert_called_once_with(expected_file_key) logic.s3.get_file_url.assert_called_once_with(expected_file_key) assert not logic.s3.upload_file_object.called assert not logic.generate_report.called def test_generate_initial_data_report_upc_report_type(default_task_data): """Should write UPC header if report type is UPC.""" task_data = default_task_data task_data['context'] = 'UPC123,UPC456' report_values = { 'reason': 'test lock reason', 'successful_isrcs': True, 'failed_isrcs': True, 'isrcs': { 'QA123': {'failed_territories': ['AF'], 'territories': []}, 'QA456': {'failed_territories': [], 'territories': ['AD']}} } task_data['result'] = report_values task = bulk_tasks.TaskStatus(**task_data) initial_report = logic._generate_initial_data_report( csv_data=task, import_report_type='upc') initial_report = test_utils.report_to_list( initial_report, separator='\r\n') expected_initial_report = [{'UPC': '="UPC123"'}, {'UPC': '="UPC456"'}] assert initial_report == expected_initial_report @pytest.mark.parametrize( 'report_type, expected_filename', [('isrc', '2017-12-01_Initial_Data_ISRC_1.csv'), ('upc', '2017-12-01_Initial_Data_UPC_1.csv')]) def test_create_file_name_no_task_type( default_task_data, report_type, expected_filename): """Should set 'Initial_data' if there is no task type.""" task_data = default_task_data task = bulk_tasks.TaskStatus(**task_data) task.type = '' created_filename = logic._create_file_name( task, import_report_type=report_type) assert created_filename == expected_filename def test_generate_initial_report_no_task_context( setup_bulk_status_db, seed_bulk_status_db, default_task_data, valid_correlation_id): """Should return error when there is no task.context.""" task = bulk_tasks.TaskStatus(**default_task_data) seed_bulk_status_db([task]) result = logic.generate_initial_report( task.id, import_report_type='isrc', correlation_id=valid_correlation_id) assert isinstance(result, response.Response) assert result.status == 404 assert result.message == error.TASK_WITHOUT_CONTEXT @pytest.mark.parametrize('statuses, expected_status', [ ( [bulk_tasks_const.SUCCESS], bulk_tasks_const.SUCCESS ), ( [bulk_tasks_const.SUCCESS, bulk_tasks_const.WARNING], bulk_tasks_const.SUCCESS ), ( [bulk_tasks_const.FAIL, bulk_tasks_const.FAIL], bulk_tasks_const.UPDATE_STATUS_FAIL ), ( [bulk_tasks_const.FAIL, bulk_tasks_const.WARNING], bulk_tasks_const.UPDATE_STATUS_PARTIAL_FAIL )]) def test_get_upc_status(statuses, expected_status): """Test _get_upc_status function""" assert logic._get_upc_status(statuses) == expected_status @pytest.mark.parametrize( 'isrc_result, expected_status, expected_isrc_upc_status', [ ( {field_const.SUCCESS: True}, bulk_tasks_const.SUCCESS, bulk_tasks_const.SUCCESS, ), ( {field_const.SUCCESS: False}, bulk_tasks_const.FAIL, bulk_tasks_const.FAIL, ), ( {field_const.CLAIMED_BY_ANOTHER_OWNER: ['QA']}, bulk_tasks_const.WARNING, bulk_tasks_const.WARNING, ), ( {field_const.RESOLVED_CONFLICT: True}, bulk_tasks_const.WARNING, bulk_tasks_const.SUCCESS, ), ( {field_const.LOCKED_TERRITORIES: ['QA']}, bulk_tasks_const.WARNING, bulk_tasks_const.WARNING, ), ( {field_const.SUBSTORE_CARVEOUTS: [['453', {'QA'}]]}, bulk_tasks_const.WARNING, bulk_tasks_const.WARNING, ), ( { field_const.SUCCESS: False, field_const.ERROR_MESSAGE: error.UPC_YOUTUBE_CARVED_OUT, }, bulk_tasks_const.WARNING, bulk_tasks_const.WARNING, ), ( {field_const.CLAIMED_BY_THE_SAME_LABEL: ['QA']}, bulk_tasks_const.FAIL, bulk_tasks_const.FAIL, )]) def test_get_isrc_status( isrc_result, expected_status, expected_isrc_upc_status): """Test _get_isrc_status function""" isrc_status, isrc_upc_status = logic._get_isrc_status(isrc_result) assert isrc_status == expected_status assert isrc_upc_status == expected_isrc_upc_status