import builtins import pytest from jinja2 import Undefined from marshmallow import Schema, ValidationError, fields from analytics.connectors.snowflake import AbstractSnowflakeQuery from tests.unit.queries.test_macros import strip_sql class ExampleSchema(Schema): test = fields.Str() table = fields.Str() field = fields.Str() number = fields.Int() database_name = fields.Str() schema_name = fields.Str() class ExampleSnowflakeQuery(AbstractSnowflakeQuery): @property def filename(self) -> str: return "test_template.sql" @property def query_schema(self) -> Schema: return ExampleSchema @property def default_params(self) -> dict: return {} class ExampleSnowflakeQueryUsingMacro(AbstractSnowflakeQuery): filename = "test_template_using_macro.sql" query_schema = ExampleSchema class ExampleSnowflakeQueryNonAsciiTemplate(AbstractSnowflakeQuery): filename = "test_template_non_ascii.sql" query_schema = ExampleSchema def test_invalid_user_input_raises(): with pytest.raises(ValidationError) as e: ExampleSnowflakeQuery({"test": 123}) assert e.value.messages == {"test": ["Not a valid string."]} def test_user_input_contains_original_input(): q = ExampleSnowflakeQuery({"test": "asd"}) assert q.user_params == {"test": "asd"} def test_query_params_is_schema_validated(): q = ExampleSnowflakeQuery({"number": "123"}) assert type(q.query_params["number"]) is int assert q.query_params["number"] == 123 def test_render_template_with_config(): q = ExampleSnowflakeQuery({"field": "name", "table": "employees"}) values = q.query_params.values() assert "name" in values assert "employees" in values def test_unused_params_excluded(): q = ExampleSnowflakeQuery( { "field": "name", "table": "employees", "key_not_in_schema": "value_not_in_schema", } ) values = q.query_params.values() assert "employees" in values assert "key_not_in_schema" not in q.query_params assert "value_not_in_schema" not in values def test_prepare_query_values(): q = ExampleSnowflakeQuery({"field": "test", "table": "test_table"}) query, bind_params = q.prepare_query() assert strip_sql(query) == "SELECT :field_1 FROM :table_2 ;" assert bind_params == { "field_1": "test", "table_2": "test_table", } def test_prepare_query_undefined_values(): q = ExampleSnowflakeQuery({}) query, bind_params = q.prepare_query() assert query.replace(" ", "").replace("\n", "") == "SELECT:field_1FROM:table_2;" assert bind_params == { "field_1": Undefined(name="field"), "table_2": Undefined(name="table"), } @pytest.mark.disable_mock_execute def test_execute_calls_orm_with_correct_args(mock_execute_orm): q = ExampleSnowflakeQuery({"field": "test", "table": "test_table"}) mock_execute_orm.reset_mock() q.execute() assert mock_execute_orm.call_count == 1 sql, args = mock_execute_orm.call_args[0] assert strip_sql(sql) == "SELECT :field_1 FROM :table_2 ;" assert args == { "field_1": "test", "table_2": "test_table", } def test_load_query_template_uses_utf8_encoding(monkeypatch): # Regression for the prod 500 caused by an em-dash in count.sql: the # slim-bookworm image has no locale set, so open() with no encoding= # kwarg defaults to ASCII and dies on any non-ASCII byte. Wrap the # builtin open to force ASCII whenever encoding isn't passed — that way # this test fails in any host locale if the explicit encoding="utf-8" # is dropped from load_query_template. real_open = builtins.open def ascii_default_open(*args, **kwargs): if "encoding" not in kwargs and len(args) < 4: kwargs["encoding"] = "ascii" return real_open(*args, **kwargs) monkeypatch.setattr(builtins, "open", ascii_default_open) q = ExampleSnowflakeQueryNonAsciiTemplate({"field": "x", "table": "t"}) template = q.load_query_template() assert "—" in template def test_query_using_macro(mock_macro_locations): q = ExampleSnowflakeQueryUsingMacro({"field": "test", "table": "test_table"}) query, bind_params = q.prepare_query() assert ( strip_sql(query) == "SELECT :field_1 macro_data as 'this is the macro', FROM :table_2 ;" ) assert bind_params == {"field_1": "test", "table_2": "test_table"} def test_raises_error(): class InvalidQuery(AbstractSnowflakeQuery): pass with pytest.raises(NotImplementedError): InvalidQuery() def test_raises_error_2(): class InvalidQuery2(AbstractSnowflakeQuery): @property def default_params(self) -> dict: return {} with pytest.raises(NotImplementedError): InvalidQuery2() def test_raises_error_4(): class InvalidQuery4(AbstractSnowflakeQuery): @property def default_params(self) -> dict: return {} @property def query_schema(self) -> Schema: return ExampleSchema @property def table_params(self) -> dict: return {} with pytest.raises(NotImplementedError): i = InvalidQuery4() i.filename