"""Unit tests for Snowflake connector base classes.""" import os from unittest.mock import MagicMock from unittest.mock import patch import pytest from snowflake.connector import errors from sqlalchemy.exc import DBAPIError from snowflake_connector.snowflake_conn import \ execute, fetchone, get_session, log_overflow, set_default_sessionmaker, \ SnowflakeDisconnectPlugin, SQLLoader @patch('snowflake_connector.snowflake_conn.sqlalchemy') @patch('snowflake_connector.snowflake_conn.sessionmaker') @patch('snowflake_connector.snowflake_conn._get_sf_config') def test_get_session( _get_sf_config_mock, sessionmaker_mock, alch_mock, sf_config_mock): """Test get_session method creates engine.""" _get_sf_config_mock.return_value = sf_config_mock with get_session(sf_config=sf_config_mock) as session: assert isinstance(session, MagicMock) alch_mock.create_engine.assert_called_with( 'snowflake://test_user:test_pass@test_host:443/test_db/' 'test_schema?account=test_acc&role=test_role&warehouse=test_wh') @patch('snowflake_connector.snowflake_conn._get_sessionmaker') @patch('snowflake_connector.snowflake_conn.text') def test__execute(text_mock, _get_sessionmaker_mock, sf_config_mock): """Test _execute method.""" query = 'SELECT :x;' session_mock = _get_sessionmaker_mock.return_value.return_value execute(query, params={'x': 'y'}, sf_config=sf_config_mock) text_mock.assert_called_with(query) session_mock.execute.assert_called_with(text_mock.return_value, {'x': 'y'}) @patch('snowflake_connector.snowflake_conn._get_sessionmaker') @patch('snowflake_connector.snowflake_conn.text') def test__execute_disconnect( text_mock, _get_sessionmaker_mock, sf_config_mock ): """Test _execute method with disconnect.""" query = 'SELECT :x;' calls = [MagicMock(), MagicMock()] _get_sessionmaker_mock.side_effect = calls calls[0].return_value.execute = MagicMock( side_effect=DBAPIError('', {}, orig=errors.DatabaseError())) session_mock = calls[1].return_value execute(query, params={'x': 'y'}, sf_config=sf_config_mock) session_mock.execute.assert_called_with(text_mock.return_value, {'x': 'y'}) @patch('snowflake_connector.snowflake_conn._get_sessionmaker') def test__execute_exception(_get_sessionmaker_mock, sf_config_mock): """Test _execute method in the case of unexpected exception.""" calls = [MagicMock(), MagicMock()] _get_sessionmaker_mock.side_effect = calls calls[0].return_value.execute = MagicMock( side_effect=Exception()) with pytest.raises(Exception): execute('SELECT :x;', params={'x': 'y'}, sf_config=sf_config_mock) @patch('snowflake_connector.snowflake_conn.sqlalchemy') @patch('snowflake_connector.snowflake_conn.sessionmaker') @patch('snowflake_connector.snowflake_conn._get_sf_config') @patch('snowflake_connector.snowflake_conn.text') def test_execute( text_mock, _get_sf_config_mock, sessionmaker_mock, alch_mock, sf_config_mock ): """Test execute method.""" query = 'SELECT :x;' _get_sf_config_mock.return_value = sf_config_mock execute(query, params={'x': 'y'}, sf_config=sf_config_mock) sessionmaker_mock.return_value.return_value.execute.assert_called_with( text_mock.return_value, {'x': 'y'}) @patch('snowflake_connector.snowflake_conn.sqlalchemy') @patch('snowflake_connector.snowflake_conn.sessionmaker') @patch('snowflake_connector.snowflake_conn._get_sf_config') @patch('snowflake_connector.snowflake_conn.text') def test_fetchone( text_mock, _get_sf_config_mock, sessionmaker_mock, alch_mock, sf_config_mock ): """Test fetchone method.""" query = 'SELECT :x;' _get_sf_config_mock.return_value = sf_config_mock fetchone(query, params={'x': 'y'}) sessionmaker_mock.return_value.return_value.execute.assert_called_with( text_mock.return_value, {'x': 'y'}) 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') @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.snowflake_conn.sessionmaker') @patch('snowflake_connector.snowflake_conn.sqlalchemy.create_engine') def test_set_default_sessionmaker( sessionmaker_mock, create_eng_mock, sf_config_mock): """Test setting the default sessionmaker with a sf_config object.""" set_default_sessionmaker(sf_config=sf_config_mock) sessionmaker_mock.assert_called() def test_snowflake_disconnect_plugin(): """Test SnowflakeDisconnectPlugin class.""" plugin = SnowflakeDisconnectPlugin('url', {}) class MyDialect: pass plugin.handle_dialect_kwargs(MyDialect, {}) dialect = MyDialect() assert dialect.is_disconnect(errors.DatabaseError(), None, None) assert not dialect.is_disconnect(RuntimeError(), None, None) @patch('snowflake_connector.snowflake_conn.logger') def test_log_overflow(logger_mock): """Test log_overflow callback.""" proxy_mock = MagicMock() proxy_mock._pool.overflow.return_value = 0 log_overflow(None, None, proxy_mock) assert logger_mock.warning.call_count == 0 proxy_mock._pool.overflow.return_value = -2 log_overflow(None, None, proxy_mock) assert logger_mock.warning.call_count == 0 proxy_mock._pool.overflow.return_value = 1 log_overflow(None, None, proxy_mock) logger_mock.warning.assert_called() proxy_mock._pool.overflow.side_effect = Exception() log_overflow(None, None, proxy_mock) logger_mock.exception.assert_called()