"""Unit tests for Snowflake connector base classes.""" import os from unittest.mock import MagicMock from unittest.mock import call from unittest.mock import patch from snowflake import connector from snowflake_connector.etl_connector import SnowflakeSQLExecutor, SQLLoader def test_init(monkeypatch, sf_config_mock): """Test __init__ method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) SnowflakeSQLExecutor(sf_config=sf_config_mock) connect_mock.assert_called_with( warehouse='test_wh', password='test_pass', account='test_acc', schema='test_schema', user='test_user', database='test_db', role='test_role', private_key='some_private_key', ocsp_fail_open=True, autocommit=True) def test_enter(monkeypatch, sf_config_mock): """Test __enter__ method to check if context manager implemented.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: assert isinstance(sf_exec, SnowflakeSQLExecutor) def test_enter_with_timeout(monkeypatch, sf_config_mock): """Test __enter__ method with the statement_timeout_in_seconds param.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) execute_mock = MagicMock() monkeypatch.setattr(SnowflakeSQLExecutor, 'execute', execute_mock) with SnowflakeSQLExecutor( sf_config=sf_config_mock, statement_timeout_in_seconds=1 ) as sf_exec: assert isinstance(sf_exec, SnowflakeSQLExecutor) assert execute_mock.called_once_with( 'ALTER SESSION SET STATEMENT_TIMEOUT_IN_SECONDS=1;') def test_exit(monkeypatch, sf_config_mock): """Test __exit__ method to check if context manager implemented.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock): pass connect_mock.assert_has_calls([call().close()]) def test_get_connection(monkeypatch, sf_config_mock): """Test get_connection method returns the result of call of connect().""" connect_mock = MagicMock(return_value=1) monkeypatch.setattr(connector, 'connect', connect_mock) assert SnowflakeSQLExecutor(sf_config=sf_config_mock).get_connection() == 1 def test_get_cursor_autocommit(monkeypatch, sf_config_mock): """Test get_cursor method in autocommit mode.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: with sf_exec.get_cursor() as cursor: pass cursor.assert_has_calls([call.close()]) def test_get_cursor_autocommit_disabled(monkeypatch, sf_config_mock): """Test get_cursor method in non-autocommit mode.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor( sf_config=sf_config_mock, autocommit=False) as sf_exec: with sf_exec.get_cursor() as cursor: pass cursor.assert_has_calls([call.execute('BEGIN')], any_order=True) connect_mock.assert_has_calls([call().commit()], any_order=True) def test_execute(monkeypatch, sf_config_mock): """Test execute method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sf_exec.execute('SELECT %(x)s;', params={'x': 'y'}) connect_mock.assert_has_calls([ call().cursor(), call().cursor().execute('SELECT %(x)s;', {'x': 'y'}), call().cursor().close()], any_order=True) def test_executemany(monkeypatch, sf_config_mock): """Test executemany method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sf_exec.executemany('SELECT %(x)s;', [{'x': 'y'}, {'x': 'z'}]) connect_mock.assert_has_calls([ call().cursor(), call().cursor().executemany('SELECT %(x)s;', [{'x': 'y'}, {'x': 'z'}]), call().cursor().close()], any_order=True) def test_fetchone(monkeypatch, sf_config_mock): """Test fetchone method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor( sf_config=sf_config_mock) as sf_exec: sf_exec.fetchone('SELECT %(x)s;', params={'x': 'y'}) connect_mock.assert_has_calls([ call().cursor(), call().cursor().execute('SELECT %(x)s;', {'x': 'y'}), call().cursor().fetchone(), call().cursor().close()], any_order=True) def test_fetchall(monkeypatch, sf_config_mock): """Test fetchall method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sf_exec.fetchall('SELECT %(x)s;', params={'x': 'y'}) connect_mock.assert_has_calls([ call().cursor(), call().cursor().execute('SELECT %(x)s;', {'x': 'y'}), call().cursor().fetchall(), call().cursor().close()], any_order=True) def test_fetchmany(monkeypatch, sf_config_mock): """Test fetchmany method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: gen = sf_exec.fetchmany('SELECT %(x)s;', 5, params={'x': 'y'}) next(gen) connect_mock.assert_has_calls([ call().cursor(), call().cursor().execute('SELECT %(x)s;', {'x': 'y'}), call().cursor().fetchmany(5)], any_order=True) def test_table_exists(monkeypatch, sf_config_mock): """Test table_exists method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sf_exec.table_exists('some_table') connect_mock.assert_has_calls([ call().cursor().execute( 'SELECT COUNT(*) FROM INFORMATION_SCHEMA.TABLES ' 'WHERE TABLE_NAME = %(table)s AND TABLE_SCHEMA = %(schema)s ' 'AND TABLE_CATALOG = %(db)s;', {'db': 'TEST_DB', 'schema': 'TEST_SCHEMA', 'table': 'SOME_TABLE'}), call().cursor().fetchone()]) def test_drop_table(monkeypatch, sf_config_mock): """Test drop_table method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sf_exec.drop_table('some_table') connect_mock.assert_has_calls([ call().cursor().execute( 'DROP TABLE IF EXISTS test_db.test_schema.some_table;', None)], any_order=True) def test_truncate_table(monkeypatch, sf_config_mock): """Test truncate_table method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sf_exec.truncate_table('some_table') connect_mock.assert_has_calls([ call().cursor().execute( 'TRUNCATE test_db.test_schema.some_table;', None)], any_order=True) def test_get_column_names(monkeypatch, sf_config_mock): """Test get_column_names method.""" connect_mock = MagicMock() connect_mock.return_value.cursor.return_value.fetchall.return_value = ( ('c1', 'desc'), ('c2', 'desc')) monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: column_names = sf_exec.get_column_names('some_table') connect_mock.assert_has_calls([ call().cursor().execute( 'DESC TABLE test_db.test_schema.some_table;', None)], any_order=True) assert column_names == ['c1', 'c2'] def test_swap_tables(monkeypatch, sf_config_mock): """Test swap_tables method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sf_exec.swap_tables('table1', 'table2') connect_mock.assert_has_calls([ call().cursor().execute( 'ALTER TABLE test_db.test_schema.table1 SWAP WITH ' 'test_db.test_schema.table2;', None)], any_order=True) def test_execute_query(monkeypatch, sf_config_mock): """Test execute_query method.""" execute_mock = MagicMock() validator_mock = MagicMock() validator_mock.format_identifiers.return_value = 'SQL', dict(a=1) sql_loader_mock = MagicMock() sql_loader_mock.load_query.return_value = 'SQL' monkeypatch.setattr(SnowflakeSQLExecutor, 'get_connection', MagicMock()) monkeypatch.setattr(SnowflakeSQLExecutor, 'execute', execute_mock) executor = SnowflakeSQLExecutor(sf_config_mock) monkeypatch.setattr(executor, 'validator', validator_mock) params = dict( column='my_favourite_column', value='my_beloved_value', table='my_cute_table' ) query_name = 'my_adorable_query' executor.execute_query(sql_loader_mock, query_name, params) sql_loader_mock.load_query.assert_any_call(query_name) validator_mock.format_identifiers.assert_any_call('SQL', params) executor.execute.assert_any_call('SQL', params=dict(a=1)) def test_fetchone_query(monkeypatch, sf_config_mock): """Test fetchone_query method.""" fetchone_response = ('response', ) fetchone_mock = MagicMock(return_value=fetchone_response) validator_mock = MagicMock() validator_mock.format_identifiers.return_value = 'SQL', dict(a=1) sql_loader_mock = MagicMock() sql_loader_mock.load_query.return_value = 'SQL' monkeypatch.setattr(SnowflakeSQLExecutor, 'get_connection', MagicMock()) monkeypatch.setattr(SnowflakeSQLExecutor, 'fetchone', fetchone_mock) executor = SnowflakeSQLExecutor(sf_config_mock) monkeypatch.setattr(executor, 'validator', validator_mock) params = dict( column='my_favourite_column', value='my_beloved_value', table='my_cute_table' ) query_name = 'my_adorable_query' res = executor.fetchone_query(sql_loader_mock, query_name, params) assert res == fetchone_response sql_loader_mock.load_query.assert_any_call(query_name) validator_mock.format_identifiers.assert_any_call('SQL', params) executor.fetchone.assert_any_call('SQL', params=dict(a=1)) def test_fetchall_query(monkeypatch, sf_config_mock): """Test fetchall_query method.""" fetchall_response = ('response', ) fetchall_mock = MagicMock(return_value=fetchall_response) validator_mock = MagicMock() validator_mock.format_identifiers.return_value = 'SQL', dict(a=1) sql_loader_mock = MagicMock() sql_loader_mock.load_query.return_value = 'SQL' monkeypatch.setattr(SnowflakeSQLExecutor, 'get_connection', MagicMock()) monkeypatch.setattr(SnowflakeSQLExecutor, 'fetchall', fetchall_mock) executor = SnowflakeSQLExecutor(sf_config_mock) monkeypatch.setattr(executor, 'validator', validator_mock) params = dict( column='my_favourite_column', value='my_beloved_value', table='my_cute_table' ) query_name = 'my_adorable_query' res = executor.fetchall_query(sql_loader_mock, query_name, params) assert res == fetchall_response sql_loader_mock.load_query.assert_any_call(query_name) validator_mock.format_identifiers.assert_any_call('SQL', params) executor.fetchall.assert_any_call('SQL', params=dict(a=1)) def test_sqlloader_constructor(): """Test for SQLLoader constructor.""" query_root = 'dir/path' loader = SQLLoader(query_root) assert loader.query_cash == {} assert loader.sql_files_root == os.path.realpath('dir/queries') assert loader.folder_version is None @patch.object(SQLLoader, '_get_sub_folder') def test_sqlloader_constructor_with_date(get_sub_folder_mock): """Test for SQLLoader constructor.""" query_root = 'dir/path' get_sub_folder_mock.return_value = '2019-01-01' loader = SQLLoader(query_root, date='2019-01-01') assert loader.query_cash == {} assert loader.sql_files_root == os.path.realpath('dir/queries') assert loader.folder_version == '2019-01-01' @patch('builtins.open') def test_sqlloader_load_query(open_mock): """Test getting SQL queries from disk.""" q_string = 'SELECT * FROM {somewhere};' open_mock.return_value.__enter__.return_value.read.return_value = q_string loader = SQLLoader('dir/path') query = loader.load_query('query') assert query == q_string path = os.path.realpath('dir/queries') + '/query.sql' open_mock.assert_called_once_with(path, 'r') @patch('builtins.open') def test_sqlloader_get_item_from_cache(open_mock): """Test getting SQL queries from cache.""" q_string = 'SELECT * FROM {somewhere};' loader = SQLLoader('dir/path') loader.query_cash = {'query': q_string} query = loader.load_query('query') assert query == q_string open_mock.assert_not_called() @patch('snowflake_connector.etl_connector.os') def test_get_sub_folder_if_folder_exists(mock_os): """Test for SQLLoader _get_sub_folder method.""" mock_os.listdir.return_value = ['test', '2019-01-01'] mock_os.path.isdir.return_value = True query_root = 'dir/path' loader = SQLLoader(query_root, date='2019-01-01') assert loader.query_cash == {} assert loader.folder_version == '2019-01-01' @patch('snowflake_connector.etl_connector.os') def test_get_sub_folder_there_is_no_folder_with_date(mock_os): """Test for SQLLoader _get_sub_folder method.""" mock_os.listdir.return_value = ['test', '2019-01-01'] mock_os.path.isdir.return_value = False query_root = 'dir/path' loader = SQLLoader(query_root, date='2019-01-01') assert loader.query_cash == {} assert loader.folder_version is None @patch.object(SQLLoader, '_get_sub_folder') def test_get_query_path_without_date(get_sub_folder_mock): """Test for SQLLoader constructor.""" query_root = 'dir/path' loader = SQLLoader(query_root) result = loader._get_query_path('query') assert result == os.path.realpath('dir/queries') + '/query.sql' get_sub_folder_mock.assert_not_called() @patch.object(SQLLoader, '_get_sub_folder') @patch('snowflake_connector.etl_connector.os') def test_get_query_path_with_date(mock_os, get_sub_folder_mock): """Test for SQLLoader constructor.""" mock_os.path.isfile.return_value = True mock_os.path.realpath.return_value = os.path.realpath('dir') query_root = 'dir/path' get_sub_folder_mock.return_value = '2019-01-01' loader = SQLLoader(query_root, date='2019-01-01') result = loader._get_query_path('query') assert result == os.path.realpath('dir/queries') + '/2019-01-01/query.sql'