"""Test snowflake_utils.py.""" from unittest.mock import MagicMock, mock_open, patch from cryptography.hazmat.primitives.asymmetric import rsa import helpers.snowflake_utils as snowflake_utils @patch('helpers.snowflake_utils.config') @patch('helpers.snowflake_utils.serialization.load_pem_private_key') def test_get_sf_private_key_from_secrets_manager(mock_load_key, mock_config): """Test get_sf_private_key with valid key from secrets manager.""" mock_config.secrets_manager_client_lambda.get_cred.side_effect = ['mock_key', 'mock_passphrase'] mock_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) mock_load_key.return_value = mock_key result = snowflake_utils.get_sf_private_key() mock_load_key.assert_called_once_with( b'mock_key', password=b'mock_passphrase', backend=snowflake_utils.default_backend() ) assert result is not None @patch('helpers.snowflake_utils.config') @patch('helpers.snowflake_utils.os') @patch('helpers.snowflake_utils.serialization.load_pem_private_key') def test_get_sf_private_key_from_local_dev(mock_load_key, mock_os, mock_config): """Test get_sf_private_key in dev environment with local key file.""" mock_config.secrets_manager_client_lambda.get_cred.side_effect = [ None, None] mock_config.ENVIRONMENT = 'dev' mock_os.environ.get.side_effect = lambda key, default=None: { 'SNOWFLAKE_PRIVATE_KEY_PATH': '/mock/path/to/key.p8', 'SNOWFLAKE_KEY_PASSPHRASE': 'mock_passphrase' }.get(key, default) mock_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) mock_load_key.return_value = mock_key with patch('builtins.open', mock_open(read_data=b'mock_key_data')): result = snowflake_utils.get_sf_private_key() mock_load_key.assert_called_once_with( b'mock_key_data', password=b'mock_passphrase', backend=snowflake_utils.default_backend() ) assert result is not None @patch('helpers.snowflake_utils.config') def test_get_sf_private_key_no_key(mock_config): """Test get_sf_private_key when no key is available.""" mock_config.secrets_manager_client_lambda.get_cred.return_value = None mock_config.ENVIRONMENT = 'prod' result = snowflake_utils.get_sf_private_key() assert result is None @patch('helpers.snowflake_utils.snowflake_connect') @patch('helpers.snowflake_utils.get_sf_private_key') @patch('helpers.snowflake_utils.config') def test_execute_snowflake_query_success(mock_config, mock_get_key, mock_connect): """Test execute_snowflake_query with successful query execution.""" # Mock configuration and private key mock_config.snowflake_db_config = {'db': 'test_db', 'schema': 'test_schema'} mock_get_key.return_value = 'mock_private_key' # Mock Snowflake connection and cursor mock_cursor = MagicMock() mock_cursor.execute.return_value.fetchmany.return_value = [{'col1': 'val1', 'col2': 'val2'}] mock_connection = MagicMock() mock_connection.cursor.return_value.__enter__.return_value = mock_cursor mock_connect.return_value.__enter__.return_value = mock_connection # Call the function query = 'SELECT * FROM {db}.{schema}.table WHERE id = %(id)s' params = {'id': 1} result = snowflake_utils.execute_snowflake_query(query, params) # Assertions mock_connect.assert_called_once_with( **mock_config.snowflake_db_config, private_key='mock_private_key' ) mock_cursor.execute.assert_called_once_with( 'SELECT * FROM test_db.test_schema.table WHERE id = %(id)s', params ) assert result == [{'col1': 'val1', 'col2': 'val2'}]