from unittest.mock import patch import pytest from jinja2 import Undefined from marshmallow import Schema, ValidationError, fields from playlist.connectors.snowflake import BASE_DIR, SnowflakeQuery class ExampleSchema(Schema): test = fields.Str() table = fields.Str() field = fields.Str() number = fields.Int() class ExampleSnowflakeQuery(SnowflakeQuery): @property def filename(self) -> str: return "test_template.sql" @property def query_schema(self) -> Schema: return ExampleSchema @property def default_params(self) -> dict: return {} @property def table_params(self) -> dict: return {} class ExampleSnowflakeQueryUsingMacro(SnowflakeQuery): filename = "test_template_using_macro.sql" query_schema = ExampleSchema def test_default_query_params_contains_db_config(): q = ExampleSnowflakeQuery({}) assert q.query_params["database_name"] == "mock_db" assert q.query_params["schema_name"] == "mock_schema" 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 "mock_db" in values assert "mock_schema" in 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 ( query.replace(" ", "").replace("\n", "") == "SELECT:field_1FROM:database_name_2.:schema_name_3.:table_4;" ) assert bind_params == { "field_1": "test", "database_name_2": "mock_db", "schema_name_3": "mock_schema", "table_4": "test_table", } def test_prepare_query_undefined_values(): q = ExampleSnowflakeQuery({}) query, bind_params = q.prepare_query() assert ( query.replace(" ", "").replace("\n", "") == "SELECT:field_1FROM:database_name_2.:schema_name_3.:table_4;" ) assert bind_params == { "field_1": Undefined(name="field"), "database_name_2": "mock_db", "schema_name_3": "mock_schema", "table_4": 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"}) q.execute() mock_execute_orm.reset() assert mock_execute_orm.call_count == 1 assert mock_execute_orm.call_args[0] == ( "SELECT\n :field_1\nFROM\n :database_name_2.:schema_name_3.:table_4\n;", { "field_1": "test", "database_name_2": "mock_db", "schema_name_3": "mock_schema", "table_4": "test_table", }, ) def test_query_using_macro(mock_macro_locations): q = ExampleSnowflakeQueryUsingMacro({"field": "test", "table": "test_table"}) query, bind_params = q.prepare_query() assert ( query.replace(" ", "").replace("\n", "") == "SELECT:field_1,macro_data as 'this is the macro',FROM:database_name_2.:schema_name_3.:table_4;" ) assert bind_params == { "field_1": "test", "database_name_2": "mock_db", "schema_name_3": "mock_schema", "table_4": "test_table", } def test_raises_error(): class InvalidQuery(SnowflakeQuery): pass with pytest.raises(NotImplementedError): InvalidQuery() def test_raises_error_2(): class InvalidQuery2(SnowflakeQuery): @property def default_params(self) -> dict: return {} with pytest.raises(NotImplementedError): InvalidQuery2() def test_raises_error_4(): class InvalidQuery4(SnowflakeQuery): @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