"""Unit tests for extra_data utils.""" import pytest from src.utils import extra_data @pytest.mark.parametrize( 'key_fields', (None, [], ()) ) def test_get_dict_values_by_keys_raises(key_fields): """Test _get_dict_values_by_keys function raises exception.""" with pytest.raises(ValueError): extra_data._get_dict_values_by_keys({}, key_fields) @pytest.mark.parametrize( 'record, key_fields, expected_result', ( ({'k1': 'v1', 'k2': 'v2'}, ['k1', 'k2'], ('v1', 'v2')), ({'k1': 'v1', 'k2': 'v2'}, ['k1'], 'v1'), ) ) def test_get_dict_values_by_keys(record, key_fields, expected_result): """Test _get_dict_values_by_keys function returns expected result.""" result = extra_data._get_dict_values_by_keys(record, key_fields) assert result == expected_result @pytest.mark.parametrize( ('records, new_data, key_fields, updated_records,' 'expected_result, remove_not_found'), ( ([{'k1': 'v1'}, {'k1': 'v2'}], {}, ['k1'], [{'k1': 'v1'}, {'k1': 'v2'}], ['v1', 'v2'], False), ([{'k1': 'v1'}, {'k1': 'v2'}, {'k1': 'v2', 'k2': 'v3'}], {}, ['k1'], [{'k1': 'v1'}, {'k1': 'v2'}, {'k1': 'v2', 'k2': 'v3'}], ['v1', 'v2'], False), ([{'k1': 'v1'}, {'k1': 'v2'}], {'v1': {'k3': 1}}, ['k1'], [{'k1': 'v1', 'k3': 1}, {'k1': 'v2'}], ['v2'], False), ([{'k1': 'v1'}, {'k1': 'v2'}], {'v1': {'k3': 1}, 'v2': {'k4': 't'}}, ['k1'], [{'k1': 'v1', 'k3': 1}, {'k1': 'v2', 'k4': 't'}], [], False), ([{'k1': 'v1'}, {'k1': 'v2'}], {}, ['k1'], [], ['v1', 'v2'], True), ([{'k1': 'v1'}, {'k1': 'v2'}], {'v1': {'k3': 1}}, ['k1'], [{'k1': 'v1', 'k3': 1}], ['v2'], True), ) ) def test_update_dicts( records, new_data, key_fields, updated_records, expected_result, remove_not_found): """Test _update_dicts function.""" result = extra_data._update_dicts( records, new_data, key_fields, remove_not_found) assert result == expected_result assert records == updated_records @pytest.mark.parametrize( 'records, func_result, key_fields, keys, not_found, func_called', ( ([{'k1': 'v1'}, {'k1': 'v2'}], {}, ['k1'], ['v1', 'v2'], ['v1', 'v2'], 1), ([{'k1': 'v1'}, {'k1': 'v2'}], {'v1': {'k3': 1}}, ['k1'], ['v1', 'v2'], ['v2'], 1), ([{'k1': 'v1'}, {'k1': 'v2'}], {'v1': {'k3': 1}, 'v2': {'k4': 't'}}, ['k1'], ['v1', 'v2'], [], 1), ([{'k1': 'v1'}, {'k1': 'v1'}], {'v1': {'k3': 1}}, ['k1'], ['v1'], [], 1), ([], {}, [], [], [], 0), ) ) def test_get_and_update( records, func_result, key_fields, keys, not_found, func_called, mocker): """Test get_and_update function.""" remove_not_found = True mocked_func = mocker.Mock(return_value=func_result) mocked_error_log = mocker.patch('src.utils.extra_data.error_log') mocked_log_missing_records = mocked_error_log.log_missing_records mocked_update_dicts = mocker.patch('src.utils.extra_data._update_dicts') mocked_update_dicts.return_value = not_found assert records == extra_data.get_and_update( records, key_fields, mocked_func, remove_not_found) assert mocked_func.call_count == func_called assert mocked_update_dicts.call_count == func_called if func_called: assert sorted(mocked_func.call_args[0][0]) == sorted(keys) assert mocked_update_dicts.call_args[0] == ( records, func_result, key_fields, remove_not_found) not_found_called = 1 if not_found and func_called else 0 assert mocked_log_missing_records.call_count == not_found_called if not_found_called: assert mocked_log_missing_records.call_args[0] == ( mocked_func, not_found)