"""Unit tests for Snowflake validator base class.""" import pytest from snowflake_connector.validator import BaseValidator validator = BaseValidator() def test_is_valid_identifier(): """Test is_valid_identifier method.""" assert validator.is_valid_identifier('tablename') assert validator.is_valid_identifier('table_name') assert validator.is_valid_identifier('tablename1') assert not validator.is_valid_identifier('1tablename') assert not validator.is_valid_identifier(';DROP TABLE;') assert validator.is_valid_identifier('";DROP TABLE;"') def test_format_identifiers(): """Test format_identifiers method.""" sql_template_src = ( 'SELECT * FROM %(db)i.%(schema)i.%(table)i WHERE field = :field;') params = dict( db='some_db', schema='some_schema', table='some_table', field='some_field') sql, non_identifier_params = validator.format_identifiers( sql_template_src, params) assert sql == ( 'SELECT * FROM some_db.some_schema.some_table ' 'WHERE field = :field;') assert non_identifier_params == {'field': 'some_field'} params = dict( db=';DROP TABLE;', schema='some_schema', table='some_table', field='some_field') with pytest.raises(AssertionError): validator.format_identifiers(sql_template_src, params)