diff --git a/sqlmesh/core/dialect.py b/sqlmesh/core/dialect.py index e4ab522198..bcfb300ba1 100644 --- a/sqlmesh/core/dialect.py +++ b/sqlmesh/core/dialect.py @@ -734,15 +734,198 @@ def parse(self: Parser) -> t.Optional[exp.Expr]: } +_SQLMESH_META_DIALECT = "sqlmesh_meta_dialect" + + +def _holds_expression(annotation: t.Any, _visited: t.Optional[t.FrozenSet[t.Any]] = None) -> bool: + """Whether a declared field type bottoms out in a SQLGlot expression. + + Covers List[exp.Expr], Optional[Dict[str, exp.DataType]], Optional[exp.Tuple], the + nested Tuple[str, Dict[str, exp.Expr]] shape used by audits/signals, and nested + Pydantic models that themselves wrap an expression field, such as `TimeColumn` + (IncrementalByTimeRangeKind.time_column). + + Stops at `_ModelKind` subclasses without recursing into their fields: a `kind` + property's own nested properties are independently dialect-tagged via the + `ModelKind` expression node's own meta when `_props_sql` recurses into them, so + treating the `kind` field itself as "holds an expression" -- true only because some + other member of the `ModelKind` union has an expression field, e.g. + `IncrementalByTimeRangeKind.time_column` -- would route its entire subtree, + including scalar sibling properties like `forward_only`, through a dialect-specific + generator and transpile them when they shouldn't be (tsql booleans becoming + `(1 = 1)`, which silently reparses as `False`). + """ + from sqlmesh.core.model.kind import _ModelKind + + if isinstance(annotation, type): + if issubclass(annotation, exp.Expr): + return True + if issubclass(annotation, _ModelKind): + return False + visited = _visited or frozenset() + if annotation in visited: + return False + if hasattr(annotation, "model_fields"): + visited = visited | {annotation} + return any( + _holds_expression(field.annotation, visited) + for field in annotation.model_fields.values() + ) + return False + return any(_holds_expression(arg, _visited) for arg in t.get_args(annotation)) + + +@functools.lru_cache(maxsize=1) +def _meta_render_policy() -> t.Dict[str, bool]: + """Map header property name -> whether its value is warehouse SQL. + + Derived from the field declarations themselves, so it stays correct as properties + are added: expression-typed values (columns, audits, physical_properties, ...) are + the user's warehouse SQL and must render in the model's dialect, while scalar-typed + values (allow_partials, description, kind, ...) are SQLMesh's own semantics and must + stay dialect-agnostic -- transpiling those is what corrupts `allow_partials TRUE` + into tsql's unparseable `(1 = 1)`. + """ + import inspect + + from sqlmesh.core.audit.definition import ModelAudit + from sqlmesh.core.metric.definition import MetricMeta + from sqlmesh.core.model import kind as kind_module + from sqlmesh.core.model.meta import ModelMeta + + sources: t.List[t.Any] = [ModelMeta, ModelAudit, MetricMeta] + sources.extend( + obj + for name, obj in vars(kind_module).items() + if inspect.isclass(obj) and hasattr(obj, "model_fields") and name.endswith("Kind") + ) + + policy: t.Dict[str, bool] = {} + for source in sources: + for name, field in source.model_fields.items(): + policy.setdefault((field.alias or name).lower(), _holds_expression(field.annotation)) + + # `ModelMeta._pre_root_validator` (sqlmesh/core/model/meta.py) renames these two + # user-facing property names to their target field before Pydantic validation, so + # they never surface as a `Field(alias=...)` for the reflection above to find. Give + # each the render policy of the field it is renamed to. + pre_validator_aliases = { + "grain": "grains", + "table_properties": "physical_properties", + } + for alias, target in pre_validator_aliases.items(): + if target in policy: + policy[alias] = policy[target] + + return policy + + +@functools.lru_cache(maxsize=None) +def _dialect_renders_array_as_brackets(dialect_name: t.Optional[str]) -> bool: + """Whether `dialect_name`'s own generator spells an array literal as `[a, b]`. + + Checked by actually rendering a sample `exp.Array` with that dialect, rather than + inspecting `Dialect.ARRAY_SIZE_NAME` or similar generator flags, because the + generator is the single source of truth for what a dialect's array syntax looks + like and there is no single shared flag for it across dialects. This also covers + dialects (tsql, sqlite, tableau, exasol, fabric) that reuse `[`/`]` for identifier + quoting and therefore render arrays as `ARRAY(...)` instead: rewriting their + `tags`/`ignored_rules` value to `[a, b]` would not be an array literal in their + grammar at all, so it silently reparses as one bracket-quoted identifier and + corrupts the value. An unrecognized dialect name renders with the generic + generator, which itself does not use brackets, so it falls back to `False`. + """ + try: + sample = exp.Array(expressions=[exp.Literal.string("x")]) + return sample.sql(dialect=dialect_name).startswith("[") + except Exception: + return False + + def _props_sql(self: Generator, expressions: t.List[exp.Expr]) -> str: props = [] size = len(expressions) for i, prop in enumerate(expressions): + parent = prop.parent + meta_dialect = parent.meta.get(_SQLMESH_META_DIALECT) if parent else None + + def render_with_model_dialect(node: exp.Expr, **overrides: t.Any) -> str: + opts: t.Dict[str, t.Any] = { + "dialect": meta_dialect, + "pretty": self.pretty, + "identify": self.identify, + "normalize": self.normalize, + "pad": self.pad, + "indent": self._indent, + "normalize_functions": self.normalize_functions, + "leading_comma": self.leading_comma, + "max_text_width": self.max_text_width, + "comments": self.comments, + } + opts.update(overrides) + + # Keep boolean literals anywhere in the value (audit args, physical_properties, + # merge_filter, ...) as `TRUE`/`FALSE`: tsql would otherwise emit `(1 = 1)`, + # which reformats differently on the next pass. The value is transpiled with + # the model dialect anyway when it is used, e.g. in the rendered audit query. + def keep_boolean_literal(n: exp.Expr) -> exp.Expr: + if not isinstance(n, exp.Boolean): + return n + literal = exp.var("TRUE" if n.this else "FALSE") + literal.comments = n.comments + return literal + + return node.transform(keep_boolean_literal).sql(**opts) + if isinstance(prop, MacroFunc): - sql = self.indent(self.sql(prop, comment=False)) + # A macro in property position wraps user-authored arguments, so it carries + # warehouse SQL the same way `columns` or `audits` do. Clear the outer node's + # own comments (not `.this`'s, which `_macro_func_sql` already attaches) + # before rendering with the model dialect, mirroring what `comment=False` + # does for the non-dialect path below -- passing `comments=False` here + # instead would build a fresh Generator with comments globally disabled, + # silently dropping every comment in the subtree rather than just the + # redundant outer one. + if meta_dialect: + prop_for_render = prop.copy() + prop_for_render.comments = None + sql = self.indent(render_with_model_dialect(prop_for_render)) + else: + sql = self.indent(self.sql(prop, comment=False)) else: - sql = self.indent(f"{prop.name} {self.sql(prop, 'value')}") + value = prop.args.get("value") + + if ( + meta_dialect + and isinstance(value, exp.Expr) + and _meta_render_policy().get(prop.name.lower()) + ): + value_sql = render_with_model_dialect(value) + elif ( + meta_dialect + and isinstance(value, exp.Array) + and _dialect_renders_array_as_brackets(meta_dialect) + ): + # Dialect-agnostic properties (e.g. `tags`, `ignored_rules`) that hold a + # list still go through the base (dialect=None) generator, which renders + # an `exp.Array` as `ARRAY(...)`. On BigQuery `ARRAY(` is parsed as a + # subquery constructor, so a multi-element `ARRAY('a', 'b')` fails to + # reparse ("Required keyword: 'value' missing for Property"). Render it + # as a bracketed list literal instead -- but only for dialects that + # actually spell arrays that way; dialects that reuse `[`/`]` for + # identifier quoting (tsql, sqlite, ...) keep the generic `ARRAY(...)` + # form, which they parse back correctly. The elements themselves stay on + # the dialect-agnostic path (`self.expressions`, not + # `render_with_model_dialect`): these are SQLMesh's own scalar values + # (tag/rule name strings), not user warehouse SQL, so they must not be + # transpiled with the model dialect (e.g. tsql boolean literals turning + # into `(1 = 1)`). + value_sql = f"[{self.expressions(value, flat=True)}]" + else: + value_sql = self.sql(prop, "value") + + sql = self.indent(f"{prop.name} {value_sql}") if i < size - 1: sql += "," @@ -853,11 +1036,29 @@ def format_model_expressions( Returns: A string representing the formatted model. """ + + def tag_meta_dialect(expression: exp.Expr) -> exp.Expr: + """Record the model dialect on meta nodes so `_props_sql` can render the + warehouse-SQL properties (columns, audits, physical_properties, ...) with it + while the SQLMesh-owned ones stay dialect-agnostic. Tags nested ModelKind + nodes too, since kinds carry expression properties of their own such as + `time_data_type` and `unique_key`.""" + if not dialect or not is_meta_expression(expression): + return expression + + expression = expression.copy() + for node in expression.find_all(Model, Audit, Metric, ModelKind): + node.meta[_SQLMESH_META_DIALECT] = dialect + expression.meta[_SQLMESH_META_DIALECT] = dialect + return expression + if len(expressions) == 1 and is_meta_expression(expressions[0]): # Meta expressions (MODEL/AUDIT/METRIC) are SQLMesh DDL, not standard SQL, # so they must never be transpiled to the target dialect (e.g. tsql would # rewrite a boolean property like `allow_partials TRUE` to `(1 = 1)`). - return expressions[0].sql( + # Individual properties whose values *are* warehouse SQL still render with + # the model dialect -- see `_props_sql` / `_meta_render_policy`. + return tag_meta_dialect(expressions[0]).sql( pretty=True, dialect=None, normalize_functions=normalize_functions ) @@ -893,7 +1094,7 @@ def cast_to_colon(node: exp.Expr) -> exp.Expr: return ";\n\n".join( # Meta expressions (MODEL/AUDIT/METRIC) are SQLMesh DDL and must stay # dialect-agnostic; only the actual query/statement expressions transpile. - expression.sql( + tag_meta_dialect(expression).sql( pretty=True, dialect=None if is_meta_expression(expression) else dialect, normalize_functions=normalize_functions, diff --git a/tests/core/test_dialect.py b/tests/core/test_dialect.py index 142b40b31f..dc55571f21 100644 --- a/tests/core/test_dialect.py +++ b/tests/core/test_dialect.py @@ -342,6 +342,372 @@ def test_format_model_expressions(): ) +@pytest.mark.parametrize( + "dialect,audit_type,int_type", + # These dialects spell the same types differently -- fabric renders a bare DATETIME2 + # with its default precision and keeps INT, tsql does the reverse. The point of the + # test is that each keeps *its own* spelling rather than being flattened. + [("tsql", "DATETIME2", "INTEGER"), ("fabric", "DATETIME2(6)", "INT")], +) +def test_format_model_expressions_meta_render_policy(dialect: str, audit_type: str, int_type: str): + """Header properties whose values are warehouse SQL render with the model dialect, + while SQLMesh's own properties stay dialect-agnostic. + + Rendering the whole header with the dialect corrupts SQLMesh DDL (tsql turns + `allow_partials TRUE` into the unparseable `(1 = 1)`), but rendering all of it + generically discards dialect-specific values the user authored, such as the + `DATETIME2` types below. The split is derived from the field declarations, so it + covers `columns`, `audits`, `physical_properties` and the expression properties + nested inside `kind` alike. + """ + formatted = format_model_expressions( + parse( + f""" + MODEL ( + name a.b, + dialect {dialect}, + kind SCD_TYPE_2_BY_TIME ( + unique_key id, + time_data_type DATETIME2(6) + ), + allow_partials true, + description 'my description', + columns ( + ts DATETIME2(6) + ), + audits ( + my_audit(threshold := CAST('2024-01-01' AS DATETIME2)) + ), + physical_properties ( + labels = (('env', 'prod')) + ) + ); + + SELECT CAST(x AS INT) AS y FROM t + """ + ), + dialect=dialect, + ) + + assert ( + formatted + == f"""MODEL ( + name a.b, + dialect {dialect}, + kind SCD_TYPE_2_BY_TIME ( + unique_key id, + time_data_type DATETIME2(6) + ), + allow_partials TRUE, + description 'my description', + columns ( + ts DATETIME2(6) + ), + audits ( + my_audit(threshold := '2024-01-01'::{audit_type}) + ), + physical_properties ( + labels = ( + ('env', 'prod') + ) + ) +); + +SELECT + x::{int_type} AS y +FROM t""" + ) + + +@pytest.mark.parametrize( + "header", + [ + "audits (my_audit(t := CAST('2024-01-01' AS DATETIME2)))", + "audits (my_audit(flag := true))", + "kind SCD_TYPE_2_BY_COLUMN(unique_key id, columns (a, b), time_data_type DATETIME2(6))", + "physical_properties (labels = (('env', 'prod')))", + "allow_partials true, description 'my description'", + "@my_prop(cutoff := CAST('2024-01-01' AS DATETIME2))", + ], +) +def test_format_model_expressions_is_idempotent(header: str): + """Formatting an already-formatted model must be a no-op. + + Rendering a dialect-specific type with the generic generator does not merely lose + formatting, it compounds: tsql `DATETIME2` renders as `TIMESTAMP`, and tsql parses + `TIMESTAMP` as ROWVERSION (a binary type), so a second pass writes `VARBINARY`. Two + runs of `sqlmesh format` silently turned a datetime into a binary type -- and for + `time_data_type` that is the physical type of the SCD valid_from/valid_to columns. + """ + source = f"MODEL (name a.b, dialect tsql, {header});\nSELECT 1 AS x" + + once = format_model_expressions(parse(source, default_dialect="tsql"), dialect="tsql") + twice = format_model_expressions(parse(once, default_dialect="tsql"), dialect="tsql") + + assert once == twice + + +@pytest.mark.parametrize( + "dialect,column_type", + [("bigquery", "DATETIME"), ("tsql", "DATETIME2(6)")], +) +def test_format_model_expressions_preserves_column_types(dialect: str, column_type: str): + """Repeated `sqlmesh format` runs must not change a model's declared column types. + + Rendering `columns` with the generic generator rewrote them: BigQuery `DATETIME` + became `TIMESTAMP` and then `TIMESTAMPTZ`, tsql `DATETIME2` became `TIMESTAMP` and + then `VARBINARY`. + """ + expected = exp.DataType.build(column_type, dialect=dialect) + formatted = f"MODEL (name a.b, dialect {dialect}, columns (ts {column_type}));\nSELECT 1 AS ts" + + for _ in range(2): + formatted = format_model_expressions( + parse(formatted, default_dialect=dialect), dialect=dialect + ) + model = load_sql_based_model(parse(formatted, default_dialect=dialect), dialect=dialect) + assert model.columns_to_types == {"ts": expected} + + +def test_format_audit_expressions_meta_render_policy(): + """AUDIT headers have their own meta model, and get the same split: `blocking` is + SQLMesh's own boolean and must not become tsql's `(1 = 0)`, while `defaults` holds + user expressions and keeps its dialect-specific type.""" + formatted = format_model_expressions( + parse( + """ + AUDIT ( + name my_audit, + dialect tsql, + blocking false, + defaults ( + cutoff := CAST('2024-01-01' AS DATETIME2) + ) + ); + + SELECT * FROM t WHERE x > 0 + """ + ), + dialect="tsql", + ) + + assert "blocking FALSE" in formatted + assert "cutoff := '2024-01-01'::DATETIME2" in formatted + + +def test_format_model_expressions_kind_time_column_dialect(): + """Expression-bearing properties nested inside `kind` render with the model dialect, + while their scalar siblings stay dialect-agnostic. + + `time_column` is a nested Pydantic model (`TimeColumn`) wrapping an expression, not + an `exp.Expr` annotation itself, so the render-policy reflection must recurse into + nested Pydantic models to classify it as warehouse SQL. Otherwise it falls back to + generic rendering and loses dialect-specific identifier quoting: tsql's `[end]` + becomes ANSI `"end"`, even though the same identifier in the query body is correctly + kept as `[end]`. + + Regression: recursing into nested Pydantic models to fix `time_column` made + `_holds_expression` also match on `kind` itself, since *some* member of the + `ModelKind` union (`IncrementalByTimeRangeKind.time_column`) holds an expression. + That routed the entire `kind (...)` subtree through a dialect-specific generator, so tsql's + boolean-literal preprocessing rewrote `forward_only TRUE` into `forward_only (1 = 1)`. + That reparses without error, but `str_to_bool` on `Paren(EQ(1, 1)).name` (`""`) + evaluates to `False`, so the value silently flips on reload. + """ + formatted = format_model_expressions( + parse( + """ + MODEL ( + name a.b, + dialect tsql, + kind INCREMENTAL_BY_TIME_RANGE ( + time_column [end], + forward_only true + ) + ); + + SELECT 1 AS x, [end] FROM t + """, + default_dialect="tsql", + ), + dialect="tsql", + ) + + assert ( + formatted + == """MODEL ( + name a.b, + dialect tsql, + kind INCREMENTAL_BY_TIME_RANGE ( + time_column [end], + forward_only TRUE + ) +); + +SELECT + 1 AS x, + [end] +FROM t""" + ) + + model = load_sql_based_model(parse(formatted, default_dialect="tsql"), dialect="tsql") + assert model.kind.forward_only is True + + +def test_format_model_expressions_macro_property_comments_preserved_with_dialect(): + """Comments inside a macro header-property must survive formatting when the model + has a `dialect` set. + + The dialect-render path goes through `Expression.sql(dialect=...)`, which builds a + fresh `Generator` with `comments` as a constructor flag: passing `comments=False` + there disables comment rendering for the *entire* subtree, rather than just + suppressing the redundant outer-level `maybe_comment` call the way `comment=False` + does for `Generator.sql()`. That previously caused comments like `/* inline note */` + to be silently dropped whenever the model declared a `dialect`. + """ + formatted = format_model_expressions( + parse( + """ + MODEL ( + name a.b, + dialect tsql, + @my_prop(cutoff := CAST('2024-01-01' AS DATETIME2) /* inline note */) + ); + + SELECT 1 AS x + """, + default_dialect="tsql", + ), + dialect="tsql", + ) + + assert "/* inline note */" in formatted + assert ( + formatted + == """MODEL ( + name a.b, + dialect tsql, + @my_prop(cutoff := '2024-01-01'::DATETIME2 /* inline note */) +); + +SELECT + 1 AS x""" + ) + + # Idempotency: formatting an already-formatted macro property must not duplicate or + # drop the comment on a second pass. + twice = format_model_expressions(parse(formatted, default_dialect="tsql"), dialect="tsql") + assert formatted == twice + + +@pytest.mark.parametrize("dialect", ["bigquery", "duckdb", "snowflake"]) +@pytest.mark.parametrize("prop_name", ["tags", "ignored_rules"]) +def test_format_model_expressions_list_property_array_literal(dialect: str, prop_name: str): + """List-valued header properties render as `[a, b]`, not `ARRAY(a, b)`, which + BigQuery parses as a subquery and fails to load.""" + source = f"""MODEL ( + name a.b, + dialect {dialect}, + {prop_name} ['C1', 'c2'] +); +SELECT 1 AS x""" + + formatted = format_model_expressions(parse(source, default_dialect=dialect), dialect=dialect) + assert f"{prop_name} ['C1', 'c2']" in formatted + + twice = format_model_expressions(parse(formatted, default_dialect=dialect), dialect=dialect) + assert formatted == twice + + model = load_sql_based_model(parse(formatted, default_dialect=dialect), dialect=dialect) + if prop_name == "tags": + assert model.tags == ["C1", "c2"] + else: + assert model.ignored_rules == {"c1", "c2"} + + +def test_format_model_expressions_array_property_no_dialect_unchanged(): + formatted = format_model_expressions( + parse("MODEL (name a.b, tags ['C1', 'c2']); SELECT 1 AS x") + ) + + assert "tags ARRAY('C1', 'c2')" in formatted + + +@pytest.mark.parametrize("dialect", ["tsql", "sqlite"]) +@pytest.mark.parametrize("prop_name", ["tags", "ignored_rules"]) +def test_format_model_expressions_list_property_bracket_identifier_dialects( + dialect: str, prop_name: str +): + """These dialects quote identifiers with `[...]`, so a bracketed list would reload + as a single identifier. List properties must keep the `ARRAY(...)` form.""" + source = f"""MODEL ( + name a.b, + dialect {dialect}, + {prop_name} ARRAY('C1', 'c2') +); +SELECT 1 AS x""" + + formatted = format_model_expressions(parse(source, default_dialect=dialect), dialect=dialect) + assert f"{prop_name} ARRAY('C1', 'c2')" in formatted + + twice = format_model_expressions(parse(formatted, default_dialect=dialect), dialect=dialect) + assert formatted == twice + + model = load_sql_based_model(parse(formatted, default_dialect=dialect), dialect=dialect) + if prop_name == "tags": + assert model.tags == ["C1", "c2"] + else: + assert model.ignored_rules == {"c1", "c2"} + + +@pytest.mark.parametrize("dialect", ["bigquery", "duckdb", "snowflake", "postgres"]) +def test_format_model_expressions_grain_alias_render_policy(dialect: str): + """`grain` is renamed to `grains` before validation, so it must share the + `grains` render policy instead of falling back to `ARRAY(...)`.""" + source = f"""MODEL ( + name a.b, + dialect {dialect}, + grain [id, id2] +); +SELECT 1 AS x, 2 AS id, 3 AS id2""" + + formatted = format_model_expressions(parse(source, default_dialect=dialect), dialect=dialect) + assert "ARRAY(id, id2)" not in formatted + + twice = format_model_expressions(parse(formatted, default_dialect=dialect), dialect=dialect) + assert formatted == twice + + model = load_sql_based_model(parse(formatted, default_dialect=dialect), dialect=dialect) + assert {c.name for c in model.grains[0].find_all(exp.Column)} == {"id", "id2"} + + +def test_format_model_expressions_table_properties_alias_render_policy(): + """`table_properties` is renamed to `physical_properties` before validation, so it + must keep dialect-specific types such as tsql's `DATETIME2`.""" + formatted = format_model_expressions( + parse( + """ + MODEL ( + name a.b, + dialect tsql, + table_properties ( + x = CAST('2024-01-01' AS DATETIME2) + ) + ); + + SELECT 1 AS x + """, + default_dialect="tsql", + ), + dialect="tsql", + ) + + assert "x = '2024-01-01'::DATETIME2" in formatted + + twice = format_model_expressions(parse(formatted, default_dialect="tsql"), dialect="tsql") + assert formatted == twice + + def test_format_model_expressions_normalize_functions(): """Regression: formatter function-name casing behavior. diff --git a/tests/core/test_format.py b/tests/core/test_format.py index 5a44e1b381..481c5c65b5 100644 --- a/tests/core/test_format.py +++ b/tests/core/test_format.py @@ -161,3 +161,37 @@ def test_format_without_state_load(tmp_path: pathlib.Path, mocker: MockerFixture context = Context(paths=tmp_path, config=Config(project="local_only"), load_state=False) context.format(check=True) mock.assert_not_called() + + +def test_format_bigquery_header_list_properties(tmp_path: pathlib.Path): + # A BigQuery model must survive `sqlmesh format` and still load: list-valued header + # properties were rewritten to `ARRAY(...)`, which BigQuery parses as a subquery. + model_file = create_temp_file( + tmp_path, + pathlib.Path("models/model.sql"), + """MODEL ( + name test.model, + kind INCREMENTAL_BY_TIME_RANGE (time_column ds), + tags ['C1', 'c2'], + ignored_rules ['noselectstar', 'ambiguousorinvalidcolumn'], + grain [id], + partitioned_by DATE_TRUNC(ds, MONTH) +); +SELECT 1 AS id, CURRENT_DATE() AS ds""", + ) + config = Config(model_defaults=ModelDefaultsConfig(dialect="bigquery")) + + Context(paths=tmp_path, config=config).format() + + formatted = model_file.read_text(encoding="utf-8") + assert "tags ['C1', 'c2']" in formatted + assert "ignored_rules ['noselectstar', 'ambiguousorinvalidcolumn']" in formatted + assert "grain [id]" in formatted + assert "partitioned_by DATE_TRUNC(ds, MONTH)" in formatted + + context = Context(paths=tmp_path, config=config) + assert context.format(check=True) + model = context.get_model("test.model") + assert model.tags == ["C1", "c2"] + assert model.ignored_rules == {"noselectstar", "ambiguousorinvalidcolumn"} + assert [p.sql("bigquery") for p in model.partitioned_by] == ["DATE_TRUNC(`ds`, MONTH)"]