Skip to content

Commit 223286f

Browse files
committed
fix(format): preserve model column dialect
Signed-off-by: Andreas Fredhøi <andreas.fredhoi@fresio.no>
1 parent 61082a6 commit 223286f

2 files changed

Lines changed: 70 additions & 5 deletions

File tree

sqlmesh/core/dialect.py

Lines changed: 32 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -703,6 +703,8 @@ def parse(self: Parser) -> t.Optional[exp.Expr]:
703703
"METRIC": _create_parser(Metric, ["name"]),
704704
}
705705

706+
_SQLMESH_COLUMNS_DIALECT = "sqlmesh_columns_dialect"
707+
706708

707709
def _props_sql(self: Generator, expressions: t.List[exp.Expr]) -> str:
708710
props = []
@@ -712,7 +714,25 @@ def _props_sql(self: Generator, expressions: t.List[exp.Expr]) -> str:
712714
if isinstance(prop, MacroFunc):
713715
sql = self.indent(self.sql(prop, comment=False))
714716
else:
715-
sql = self.indent(f"{prop.name} {self.sql(prop, 'value')}")
717+
value = prop.args.get("value")
718+
parent = prop.parent
719+
columns_dialect = parent.meta.get(_SQLMESH_COLUMNS_DIALECT) if parent else None
720+
if prop.name.lower() == "columns" and columns_dialect and isinstance(value, exp.Expr):
721+
value_sql = value.sql(
722+
dialect=columns_dialect,
723+
pretty=self.pretty,
724+
identify=self.identify,
725+
normalize=self.normalize,
726+
pad=self.pad,
727+
indent=self._indent,
728+
normalize_functions=self.normalize_functions,
729+
leading_comma=self.leading_comma,
730+
max_text_width=self.max_text_width,
731+
comments=self.comments,
732+
)
733+
else:
734+
value_sql = self.sql(prop, "value")
735+
sql = self.indent(f"{prop.name} {value_sql}")
716736

717737
if i < size - 1:
718738
sql += ","
@@ -819,11 +839,18 @@ def format_model_expressions(
819839
Returns:
820840
A string representing the formatted model.
821841
"""
842+
843+
def expression_with_columns_dialect(expression: exp.Expr) -> exp.Expr:
844+
if isinstance(expression, Model) and dialect:
845+
expression = expression.copy()
846+
expression.meta[_SQLMESH_COLUMNS_DIALECT] = dialect
847+
return expression
848+
822849
if len(expressions) == 1 and is_meta_expression(expressions[0]):
823850
# Meta expressions (MODEL/AUDIT/METRIC) are SQLMesh DDL, not standard SQL,
824851
# so they must never be transpiled to the target dialect (e.g. tsql would
825852
# rewrite a boolean property like `allow_partials TRUE` to `(1 = 1)`).
826-
return expressions[0].sql(
853+
return expression_with_columns_dialect(expressions[0]).sql(
827854
pretty=True, dialect=None, normalize_functions=normalize_functions
828855
)
829856

@@ -857,9 +884,9 @@ def cast_to_colon(node: exp.Expr) -> exp.Expr:
857884
expressions = new_expressions
858885

859886
return ";\n\n".join(
860-
# Meta expressions (MODEL/AUDIT/METRIC) are SQLMesh DDL and must stay
861-
# dialect-agnostic; only the actual query/statement expressions transpile.
862-
expression.sql(
887+
# Meta expressions (MODEL/AUDIT/METRIC) stay dialect-agnostic, except for
888+
# MODEL column types. Actual query/statement expressions are transpiled.
889+
expression_with_columns_dialect(expression).sql(
863890
pretty=True,
864891
dialect=None if is_meta_expression(expression) else dialect,
865892
normalize_functions=normalize_functions,

tests/core/test_dialect.py

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -342,6 +342,44 @@ def test_format_model_expressions():
342342
)
343343

344344

345+
@pytest.mark.parametrize("dialect", ["tsql", "fabric"])
346+
def test_format_model_columns_uses_model_dialect(dialect: str):
347+
formatted = format_model_expressions(
348+
parse(
349+
f"""
350+
MODEL (
351+
name test_model,
352+
dialect {dialect},
353+
description 'my description',
354+
formatting false,
355+
columns (
356+
_dwh_load_datetime_utc DATETIME2(6)
357+
)
358+
);
359+
360+
SELECT 1 AS id
361+
"""
362+
),
363+
dialect=dialect,
364+
)
365+
366+
assert (
367+
formatted
368+
== f"""MODEL (
369+
name test_model,
370+
dialect {dialect},
371+
description 'my description',
372+
formatting FALSE,
373+
columns (
374+
_dwh_load_datetime_utc DATETIME2(6)
375+
)
376+
);
377+
378+
SELECT
379+
1 AS id"""
380+
)
381+
382+
345383
def test_format_model_expressions_normalize_functions():
346384
"""Regression: formatter function-name casing behavior.
347385

0 commit comments

Comments
 (0)