"""Unit tests for Snowflake validator base class.""" from unittest.mock import MagicMock import pytest from snowflake import connector from snowflake_connector.etl_connector import SnowflakeSQLExecutor def test_is_valid_identifier(monkeypatch, sf_config_mock): """Test is_valid_identifier method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: assert sf_exec.validator.is_valid_identifier('tablename') assert sf_exec.validator.is_valid_identifier('table_name') assert sf_exec.validator.is_valid_identifier('tablename1') assert not sf_exec.validator.is_valid_identifier('1tablename') assert not sf_exec.validator.is_valid_identifier(';DROP TABLE;') assert sf_exec.validator.is_valid_identifier('";DROP TABLE;"') def test_format_identifiers(monkeypatch, sf_config_mock): """Test format_identifiers method.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) sql_template_src = ( 'SELECT * FROM %(db)i.%(schema)i.%(table)i WHERE field = %(field)s;') params = dict( db='some_db', schema='some_schema', table='some_table', field='some_field') with SnowflakeSQLExecutor(sf_config=sf_config_mock) as sf_exec: sql, non_identifier_params = sf_exec.validator.format_identifiers( sql_template_src, params) assert sql == ( 'SELECT * FROM some_db.some_schema.some_table ' 'WHERE field = %(field)s;') assert non_identifier_params == {'field': 'some_field'} assert params == dict( # params dict is kept intact db='some_db', schema='some_schema', table='some_table', field='some_field' ) params = dict( db=';DROP TABLE;', schema='some_schema', table='some_table', field='some_field') with pytest.raises(Exception): sf_exec.validator.format_identifiers(sql_template_src, params)