import sqlglot def _is_valid_sql(sql: str) -> bool: """Checks if the provided SQL query is syntactically valid.""" try: sqlglot.parse(sql) return True except sqlglot.ParseError: return False def pytest_generate_tests(metafunc): if "query_func" in metafunc.fixturenames: funcs = getattr(metafunc.cls, "QUERY_FUNCS", []) metafunc.parametrize("query_func", funcs) class BaseSqlQueryTest: """Base class for SQL query tests, both for producer and consumer.""" QUERY_FUNCS = [] MODEL = None def test_sql_queries_valid_sql(self, query_func): result = query_func() assert _is_valid_sql(result), f"Invalid SQL: {result}" def test_sql_queries_all_final_columns_in_model(self, query_func): result = query_func() expr = sqlglot.parse_one(result) select = expr.find(sqlglot.expressions.Select) final_columns = [e.alias_or_name for e in select.expressions] model_fields_aliases = {f.alias for f in self.MODEL.model_fields.values()} assert all( col in model_fields_aliases for col in final_columns ), f"All output columns must be in the model: {final_columns}"