"""Geocoding flow tests.""" from unittest import mock from unittest.mock import call from unittest.mock import MagicMock from unittest.mock import patch from dim_refresh_etl.conf.config import SF_CONFIG from dim_refresh_etl.flows.geocoding import settings from dim_refresh_etl.flows.geocoding import tasks SF_CONFIG_MOCK = { 'db': 'db', 'schema': 'schema', 'role': 'role', 'warehouse': 'warehouse', } @patch('dim_refresh_etl.flows.geocoding.tasks.geocoder') def test_write_lat_long_to_s3_file(mock_geocoder): """Test _write_lat_long_to_s3_file method.""" mock_coordinate = MagicMock() mock_coordinate.latlng = [23.4634, 23.6666] mock_geocoder.google.return_value = mock_coordinate mock_file_object = MagicMock() mock_file_object.write.return_value = None tasks._write_lat_long_to_s3_file( mock_file_object, 'US', '11375') mock_file_object.write.assert_any_call( '11375\tUS\t23.4634\t23.6666\n') @patch('dim_refresh_etl.flows.geocoding.tasks.smart_open') @patch('dim_refresh_etl.flows.geocoding.tasks.SnowflakeSQLExecutor') @patch('dim_refresh_etl.flows.geocoding.tasks.settings') @patch('dim_refresh_etl.flows.geocoding.tasks._write_lat_long_to_s3_file') def test_write_country_zip_code_lat_long_map( mock_write_function, settings_mock, executor_class_mock, mock_smart_open): """Test write_country_zip_code_lat_long_map method.""" settings_mock.SELECT_SQL = 'select' executor_mock = MagicMock() cursor_mock = MagicMock() executor_class_mock.return_value.__enter__. \ return_value = executor_mock executor_mock.get_cursor.return_value.__enter__.return_value = cursor_mock executor_mock.validator.format_identifiers.return_value = ('select', '') with patch.dict(SF_CONFIG, SF_CONFIG_MOCK, clean=True): tasks.write_country_zip_code_lat_long_map(MagicMock(), {}) cursor_mock.execute.assert_any_call('select') assert not mock_write_function.called @patch('dim_refresh_etl.flows.geocoding.tasks.SnowflakeSQLExecutor') @patch('dim_refresh_etl.flows.geocoding.tasks.settings') def test_update_lat_long_on_dim_zip(settings_mock, executor_class_mock): """Test update_lat_long_on_dim_zip method.""" settings_mock.CREATE_TEMP_TABLE_SQL = 'create' settings_mock.LOAD_TEMP_TABLE_SQL = 'load' settings_mock.UPDATE_SQL = 'update' execute_mock = MagicMock(return_value=None) executor_mock = MagicMock() executor_class_mock.return_value.__enter__. \ return_value = executor_mock executor_mock.execute = executor_mock executor_mock.get_connection.return_value = MagicMock() executor_mock.execute = execute_mock executor_mock.validator.format_identifiers.side_effect = [ ('create', ''), ('load', ''), ('update', '')] with patch.dict(SF_CONFIG, SF_CONFIG_MOCK, clean=True): tasks.update_lat_long_on_dim_zip(MagicMock(), None) execute_mock.assert_has_calls([ call('create'), call('load', mock.ANY), call('update')]) def _test_get_zip_country_code_sql(): """Test _get_zip_country_code_sql method.""" sql = settings.SELECT_SQL new_sql = tasks._get_zip_country_code_sql('NL') # noqa assert new_sql == '{}{}{}'.format( sql, (" AND country_code = 'NL' " "AND requested = 'N' " 'ORDER BY date_created DESC LIMIT '), settings.REQUEST_LIMIT)