"""Tests for query-generators related to distribution_fee ETL.""" from argparse import ArgumentTypeError from unittest.mock import Mock from unittest.mock import patch from pytest import raises from flows.distribution_fee import queries def test_get_delete_client_amount_sql(): """Test get_delete_client_amount_sql function.""" upcs = ['123', '234', '345', '456', '567'] upcs_expected = [['123', '234', '345'], ['456', '567']] queries.get_delete_client_amount_sql.batch_size = 3 results_raw = queries.get_delete_client_amount_sql(upcs) results = list(results_raw) assert len(results) == len(upcs_expected) for i, result in enumerate(results): for upc in upcs_expected[i]: assert upc in result def test_get_insert_client_amount_sql(): """Test get_insert_client_amount_sql function.""" upcs = ['123', '234', '345', '456', '567'] upcs_expected = [['123', '234', '345'], ['456', '567']] params = { 'table': 'mytable', 'column': 'mycolumn', 'fee_dms_table': 'fee_dms_table', 'fee_ter_table': 'fee_ter_table', 'fee_reg_table': 'fee_reg_table', 'upcs': upcs} queries.get_insert_client_amount_sql.batch_size = 3 results_raw = queries.get_insert_client_amount_sql(**params) results = list(results_raw) assert len(results) == len(upcs_expected) for i, result in enumerate(results): for upc in upcs_expected[i]: assert upc in result def test_get_delete_distribution_table_data_sql(): """Test get_delete_distribution_table_data_sql function.""" upcs = ['123', '234', '345', '456', '567'] upcs_expected = [['123', '234', '345'], ['456', '567']] queries.get_delete_distribution_table_data_sql.batch_size = 3 results_raw = queries.get_delete_distribution_table_data_sql(upcs) results = list(results_raw) assert len(results) == len(upcs_expected) for i, result in enumerate(results): for upc in upcs_expected[i]: assert upc in result def test_get_insert_distribution_table_data_sql(): """Test get_insert_distribution_table_data_sql function.""" results_dms = queries.get_insert_distribution_table_data_sql( 'temp_fee_dms_table', 'dms') results_territory = queries.get_insert_distribution_table_data_sql( 'temp_fee_territory_table', 'territory') results_regular = queries.get_insert_distribution_table_data_sql( 'temp_fee_regular_table', 'regular') assert 'dms' in results_dms assert 'country_id' in results_dms assert 'territory' in results_territory assert 'regular' in results_regular assert 'NULL' in results_regular @patch('flows.distribution_fee.queries.SELECT_FEE_REGULAR') def test_get_select_contract_fee_sql_regular(select_query): """Test get_select_contract_fee_sql function for regular contracts.""" formatted_query_0 = Mock() formatted_query_1 = Mock() select_query.format.side_effect = [formatted_query_0, formatted_query_1] contract_ids_0 = [12, 23, 34] contract_ids_1 = [45, 56] contract_ids = contract_ids_0 + contract_ids_1 queries.get_select_contract_fee_sql.batch_size = 3 results_raw = queries.get_select_contract_fee_sql(contract_ids, 'regular') results = list(results_raw) assert len(results) == 2 assert results == [formatted_query_0, formatted_query_1] format_calls = select_query.format.call_args_list for contract_id in contract_ids_0: assert str(contract_id) in format_calls[0][1]['vendor_contract_ids'] for contract_id in contract_ids_1: assert str(contract_id) in format_calls[1][1]['vendor_contract_ids'] @patch('flows.distribution_fee.queries.SELECT_FEE_TERRITORY') def test_get_select_contract_fee_sql_territory(select_query): """Test get_select_contract_fee_sql function for territory contracts.""" formatted_query_0 = Mock() formatted_query_1 = Mock() select_query.format.side_effect = [formatted_query_0, formatted_query_1] contract_ids_0 = [12, 23, 34] contract_ids_1 = [45, 56] contract_ids = contract_ids_0 + contract_ids_1 queries.get_select_contract_fee_sql.batch_size = 3 results_raw = queries.get_select_contract_fee_sql( contract_ids, 'territory') results = list(results_raw) assert len(results) == 2 assert results == [formatted_query_0, formatted_query_1] format_calls = select_query.format.call_args_list for contract_id in contract_ids_0: assert str(contract_id) in format_calls[0][1]['vendor_contract_ids'] for contract_id in contract_ids_1: assert str(contract_id) in format_calls[1][1]['vendor_contract_ids'] def test_get_select_contract_fee_sql_exception(): """Test get_select_contract_fee_sql function with bad inputs.""" contract_ids = [12, 23, 34, 'bad'] with raises(ArgumentTypeError): results_raw = queries.get_select_contract_fee_sql( contract_ids, 'territory') list(results_raw) # trigger iteration and exception