"""Test for tasks module.""" from datetime import date import itertools from unittest.mock import Mock from unittest.mock import patch from flows.distribution_fee import status from flows.distribution_fee import tasks @patch('flows.distribution_fee.tasks.log') @patch('flows.distribution_fee.tasks.datastore') def test_create_temp_table(datastore_mock, log_mock): """Test create_temp_table function.""" table_name = 'test_table_for_this_unit_test' create_statement = 'create table {table_name}' params = { 'activity': Mock(), 'correlation_id': 'cid1234', 'table_name': table_name, 'create_statement': create_statement} result = tasks.create_temp_table(**params) expected = {'name': table_name} assert result == expected db_call = datastore_mock.execute.call_args_list[0][0] assert table_name in db_call[0] @patch('flows.distribution_fee.tasks.art_relations') @patch('flows.distribution_fee.tasks.datastore') @patch('flows.distribution_fee.tasks.log') @patch('flows.distribution_fee.tasks.queries.get_select_contract_fee_sql') @patch('flows.distribution_fee.util.datastore') def test_load_dist_fee_territory( util_datastore, get_select_sql, log, task_datastore, art_relations, vendor_contracts, distribution_fees_territory, distribution_fees_territory_split): """Test load_dist_fee_territory task function.""" ds_cursor = Mock() ds_cursor.fetchall.return_value = vendor_contracts util_datastore.query.return_value = ds_cursor art_relations.query.side_effect = [ Mock(fetchall=Mock(return_value=distribution_fees_territory[0:2])), Mock(fetchall=Mock(return_value=distribution_fees_territory[2:4])), Mock(fetchall=Mock(return_value=distribution_fees_territory[4:6]))] select_queries = ['query 1', 'query 2', 'query 3'] get_select_sql.return_value = select_queries fee_table = 'test_fees' contracts_table = 'test_contracts' correlation_id = '0123-4567-8910-1112' tasks.load_dist_fee_territory( None, contracts_table, correlation_id, fee_table) # assert only contract IDs from temp table are used for query generation for select_query in select_queries: art_relations.query.assert_any_call(select_query) # assert only contract IDs from temp table are used for query generation query_get = get_select_sql.call_args_list[0][0] assert {contract[0] for contract in vendor_contracts} == set(query_get[0]) assert 'territory' == query_get[1] # assert expected data set to be inserted into the datastore temp table insert_query, insert_sequence = \ task_datastore.executemany.call_args_list[0][0] assert fee_table in insert_query assert len(insert_sequence) == len(distribution_fees_territory_split) for insert_row in insert_sequence: assert tuple(insert_row) in distribution_fees_territory_split log.update_status.assert_called_once_with( correlation_id, status.DIST_FEE_TERRITORY_LOADED) @patch('flows.distribution_fee.tasks.util') @patch('flows.distribution_fee.tasks.art_relations') @patch('flows.distribution_fee.tasks.datastore') @patch('flows.distribution_fee.tasks.log') def test_load_vendor_contract( log, datastore, art_relations, util, upc_vendor_mapping, vendor_contract_ids): """Test load_vendor_contract function.""" correlation_id = 'coid-1234-5678' table_name = 'test_table_for_this_unit_test' call_query = 'the call statement' insert_query = 'the insert into {table_name} statement' # Adding 'abc' UPC to test case when that UPC/Vendor ID does not have a # contract ('abc' or 'cba' key does not live in vendor_contract_ids). upc_vendor_mapping_fixture = upc_vendor_mapping.copy() upc_vendor_mapping_fixture['abc'] = 'cba' def get_side_effect(upcs): results = [] for upc, vendor_id in upc_vendor_mapping_fixture.items(): if upc in upcs: results.append({upc: vendor_id}) return results util.get_upc_vendor_id.side_effect = get_side_effect def query_side_effect(call_query, args): vendor_id, *_ = args sp_call_mock = Mock() sp_call_mock.fetchone.side_effect = \ lambda: vendor_contract_ids.get(vendor_id, None) return sp_call_mock art_relations.query.side_effect = query_side_effect params = { 'activity': Mock(), 'call_query': call_query, 'correlation_id': correlation_id, 'insert_query': insert_query, 'table_name': table_name, 'upcs': ['123', '234', '345', 'abc', 'not_exist']} tasks.load_vendor_contract(**params) called_upcs = set() for db_call_args, _ in datastore.executemany.call_args_list: assert table_name in db_call_args[0] rows = db_call_args[1] for vendor_contract_id, vendor_id, upc in rows: assert vendor_contract_ids[vendor_id][0] == vendor_contract_id called_upcs.add(upc) # all specified and existing upcs were called assert not {'123', '234', '345'} - called_upcs log.update_status.assert_called_once_with( correlation_id, status.LOADED_CONTRACT_IDS) @patch('flows.distribution_fee.tasks.art_relations') @patch('flows.distribution_fee.tasks.datastore') @patch('flows.distribution_fee.tasks.log') @patch('flows.distribution_fee.tasks.queries.get_select_contract_fee_sql') @patch('flows.distribution_fee.util.datastore') def test_load_dist_fee_regular( util_datastore, get_select_sql, log, task_datastore, art_relations, vendor_contracts, distribution_fees_regular, distribution_fees_regular_split): """Test load_dist_fee_regular task function.""" ds_cursor = Mock() ds_cursor.fetchall.return_value = vendor_contracts util_datastore.query.return_value = ds_cursor art_relations.query.side_effect = [ Mock(fetchall=Mock(return_value=distribution_fees_regular[0:2])), Mock(fetchall=Mock(return_value=distribution_fees_regular[2:4])), Mock(fetchall=Mock(return_value=distribution_fees_regular[4:6]))] select_queries = ['query 1', 'query 2', 'query 3'] get_select_sql.return_value = select_queries fee_table = 'test_fees' contracts_table = 'test_contracts' correlation_id = '0123-4567-8910-1112' tasks.load_dist_fee_regular( None, contracts_table, correlation_id, fee_table) for select_query in select_queries: art_relations.query.assert_any_call(select_query) # assert only contract IDs from temp table are used for query generation query_get = get_select_sql.call_args_list[0][0] assert {contract[0] for contract in vendor_contracts} == set(query_get[0]) assert 'regular' == query_get[1] # assert expected data set to be inserted into the datastore temp table insert_query, insert_sequence = \ task_datastore.executemany.call_args_list[0][0] assert fee_table in insert_query assert len(insert_sequence) == len(distribution_fees_regular_split) for insert_row in insert_sequence: assert tuple(insert_row) in distribution_fees_regular_split log.update_status.assert_called_once_with( correlation_id, status.DIST_FEE_REGULAR_LOADED) @patch('flows.distribution_fee.tasks.log') @patch('flows.distribution_fee.tasks.datastore') def test_drop_temp_table(datastore_mock, log_mock): """Test drop_temp_table function.""" table_name = 'test_table_for_this_unit_test' drop_statement = 'drop if exists table {table_name}' params = { 'activity': Mock(), 'correlation_id': 'cid1234', 'table_name': table_name, 'drop_statement': drop_statement} tasks.drop_temp_table(**params) db_call = datastore_mock.execute.call_args_list[0][0] assert table_name in db_call[0] @patch('flows.distribution_fee.tasks.log') @patch('flows.distribution_fee.tasks.feed_status') def test_set_dynamo_status(feed_status_mock, log_mock): """Test set_dynamo_status task.""" params = { 'activity': Mock(), 'correlation_id': 'cid1234'} tasks.set_dynamo_status(**params) assert feed_status_mock.set_overall_status.called @patch('flows.distribution_fee.tasks.date') @patch('flows.distribution_fee.tasks.etl_util') def test_send_sns_notification(etl_util_mock, date_mock): """Test send_sns_notification function.""" activity = Mock() correlation_id = '1234-5678-9101-1120.1' sns_correlation_id = '1234-5678-9101-1120.1.1' report_date = date(2017, 3, 15) upcs = ['123', '456'] date_mock.today.return_value = report_date tasks.send_sns_notification(activity, correlation_id, upcs) etl_util_mock.send_sns_notification.assert_called_with( sns_correlation_id, report_date.strftime('%Y-%m-%d'), upcs) @patch('flows.distribution_fee.tasks.log') @patch('flows.distribution_fee.tasks.datastore') @patch('flows.distribution_fee.tasks.config') def test_calculate_client_amount( config_mock, datastore_mock, log_mock, database_context): """Test calculate_client_amount function.""" tasks.queries.get_delete_client_amount_sql.batch_size = 3 tasks.queries.get_insert_client_amount_sql.batch_size = 3 upcs = ('123', '234', '345', '456', '567') upcs_batches = (('123', '234', '345'), ('456', '567')) params = { 'activity': Mock(), 'correlation_id': 'cid1234', 'fee_dms_table': 'fee_dms_table_name', 'fee_ter_table': 'fee_ter_table_name', 'fee_reg_table': 'fee_reg_table_name', 'upcs': upcs} config_mock.REVENUE_TABLE_COLUMNS = { 'foo_table': 'foo_column', 'bar_table': 'bar_column', 'baz_table': 'baz_column'} datastore_mock.context = database_context tasks.calculate_client_amount(**params) # These params are string searches for the db cursor execute calls. Rather # than search for exact SQL strings this tests deletes and inserts params. execute_params = list( itertools.product( ('INSERT INTO',), config_mock.REVENUE_TABLE_COLUMNS.items(), upcs_batches)) execute_params += list( itertools.product( ('DELETE FROM',), (('client_amount', 'upc'),), upcs_batches)) execute_calls = database_context._cursor.execute.call_args_list for action, (table, column), upcs in execute_params: call_found = False # every single iteration of params should be found for call in execute_calls: sql = call[0][0] call_found = all(( action in sql, table in sql, column in sql, all(upc in sql for upc in upcs))) if call_found: break assert call_found, 'Expected SQL query call parameters not found.' assert len(execute_calls) == len(execute_params) assert log_mock.update_status.called @patch('flows.distribution_fee.tasks.queries') @patch('flows.distribution_fee.tasks.log') @patch('flows.distribution_fee.tasks.datastore') def test_update_distribution_fee_table( datastore_mock, log_mock, queries_mock, database_context): """Test update_distribution_fee_table function.""" upcs = ('123', '234', '345', '456', '567') params = { 'activity': Mock(), 'correlation_id': 'cid1234', 'fee_dms_table': 'fee_dms_table_name', 'fee_ter_table': 'fee_ter_table_name', 'fee_reg_table': 'fee_reg_table_name', 'upcs': upcs} queries_mock.get_delete_distribution_table_data_sql.return_value = \ ['DELETE'] queries_mock.get_insert_distribution_table_data_sql.return_value = 'INSERT' datastore_mock.context = database_context tasks.update_distribution_fee_table(**params) execute_calls = database_context._cursor.execute.call_args_list queries_mock.get_delete_distribution_table_data_sql.\ assert_called_once_with(upcs) queries_mock.get_insert_distribution_table_data_sql.assert_any_call( 'fee_dms_table_name', 'dms') queries_mock.get_insert_distribution_table_data_sql.assert_any_call( 'fee_ter_table_name', 'territory') queries_mock.get_insert_distribution_table_data_sql.assert_any_call( 'fee_reg_table_name', 'regular') assert len(execute_calls) == 4 assert log_mock.update_status.called