Skip to content

Commit 1da53a1

Browse files
authored
fix(format): preserve dialect-specific SQL in MODEL/AUDIT/METRIC headers (#5949)
Signed-off-by: mday-io <mdaytn@gmail.com>
1 parent 5f1911d commit 1da53a1

3 files changed

Lines changed: 605 additions & 4 deletions

File tree

‎sqlmesh/core/dialect.py‎

Lines changed: 205 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -734,15 +734,198 @@ def parse(self: Parser) -> t.Optional[exp.Expr]:
734734
}
735735

736736

737+
_SQLMESH_META_DIALECT = "sqlmesh_meta_dialect"
738+
739+
740+
def _holds_expression(annotation: t.Any, _visited: t.Optional[t.FrozenSet[t.Any]] = None) -> bool:
741+
"""Whether a declared field type bottoms out in a SQLGlot expression.
742+
743+
Covers List[exp.Expr], Optional[Dict[str, exp.DataType]], Optional[exp.Tuple], the
744+
nested Tuple[str, Dict[str, exp.Expr]] shape used by audits/signals, and nested
745+
Pydantic models that themselves wrap an expression field, such as `TimeColumn`
746+
(IncrementalByTimeRangeKind.time_column).
747+
748+
Stops at `_ModelKind` subclasses without recursing into their fields: a `kind`
749+
property's own nested properties are independently dialect-tagged via the
750+
`ModelKind` expression node's own meta when `_props_sql` recurses into them, so
751+
treating the `kind` field itself as "holds an expression" -- true only because some
752+
other member of the `ModelKind` union has an expression field, e.g.
753+
`IncrementalByTimeRangeKind.time_column` -- would route its entire subtree,
754+
including scalar sibling properties like `forward_only`, through a dialect-specific
755+
generator and transpile them when they shouldn't be (tsql booleans becoming
756+
`(1 = 1)`, which silently reparses as `False`).
757+
"""
758+
from sqlmesh.core.model.kind import _ModelKind
759+
760+
if isinstance(annotation, type):
761+
if issubclass(annotation, exp.Expr):
762+
return True
763+
if issubclass(annotation, _ModelKind):
764+
return False
765+
visited = _visited or frozenset()
766+
if annotation in visited:
767+
return False
768+
if hasattr(annotation, "model_fields"):
769+
visited = visited | {annotation}
770+
return any(
771+
_holds_expression(field.annotation, visited)
772+
for field in annotation.model_fields.values()
773+
)
774+
return False
775+
return any(_holds_expression(arg, _visited) for arg in t.get_args(annotation))
776+
777+
778+
@functools.lru_cache(maxsize=1)
779+
def _meta_render_policy() -> t.Dict[str, bool]:
780+
"""Map header property name -> whether its value is warehouse SQL.
781+
782+
Derived from the field declarations themselves, so it stays correct as properties
783+
are added: expression-typed values (columns, audits, physical_properties, ...) are
784+
the user's warehouse SQL and must render in the model's dialect, while scalar-typed
785+
values (allow_partials, description, kind, ...) are SQLMesh's own semantics and must
786+
stay dialect-agnostic -- transpiling those is what corrupts `allow_partials TRUE`
787+
into tsql's unparseable `(1 = 1)`.
788+
"""
789+
import inspect
790+
791+
from sqlmesh.core.audit.definition import ModelAudit
792+
from sqlmesh.core.metric.definition import MetricMeta
793+
from sqlmesh.core.model import kind as kind_module
794+
from sqlmesh.core.model.meta import ModelMeta
795+
796+
sources: t.List[t.Any] = [ModelMeta, ModelAudit, MetricMeta]
797+
sources.extend(
798+
obj
799+
for name, obj in vars(kind_module).items()
800+
if inspect.isclass(obj) and hasattr(obj, "model_fields") and name.endswith("Kind")
801+
)
802+
803+
policy: t.Dict[str, bool] = {}
804+
for source in sources:
805+
for name, field in source.model_fields.items():
806+
policy.setdefault((field.alias or name).lower(), _holds_expression(field.annotation))
807+
808+
# `ModelMeta._pre_root_validator` (sqlmesh/core/model/meta.py) renames these two
809+
# user-facing property names to their target field before Pydantic validation, so
810+
# they never surface as a `Field(alias=...)` for the reflection above to find. Give
811+
# each the render policy of the field it is renamed to.
812+
pre_validator_aliases = {
813+
"grain": "grains",
814+
"table_properties": "physical_properties",
815+
}
816+
for alias, target in pre_validator_aliases.items():
817+
if target in policy:
818+
policy[alias] = policy[target]
819+
820+
return policy
821+
822+
823+
@functools.lru_cache(maxsize=None)
824+
def _dialect_renders_array_as_brackets(dialect_name: t.Optional[str]) -> bool:
825+
"""Whether `dialect_name`'s own generator spells an array literal as `[a, b]`.
826+
827+
Checked by actually rendering a sample `exp.Array` with that dialect, rather than
828+
inspecting `Dialect.ARRAY_SIZE_NAME` or similar generator flags, because the
829+
generator is the single source of truth for what a dialect's array syntax looks
830+
like and there is no single shared flag for it across dialects. This also covers
831+
dialects (tsql, sqlite, tableau, exasol, fabric) that reuse `[`/`]` for identifier
832+
quoting and therefore render arrays as `ARRAY(...)` instead: rewriting their
833+
`tags`/`ignored_rules` value to `[a, b]` would not be an array literal in their
834+
grammar at all, so it silently reparses as one bracket-quoted identifier and
835+
corrupts the value. An unrecognized dialect name renders with the generic
836+
generator, which itself does not use brackets, so it falls back to `False`.
837+
"""
838+
try:
839+
sample = exp.Array(expressions=[exp.Literal.string("x")])
840+
return sample.sql(dialect=dialect_name).startswith("[")
841+
except Exception:
842+
return False
843+
844+
737845
def _props_sql(self: Generator, expressions: t.List[exp.Expr]) -> str:
738846
props = []
739847
size = len(expressions)
740848

741849
for i, prop in enumerate(expressions):
850+
parent = prop.parent
851+
meta_dialect = parent.meta.get(_SQLMESH_META_DIALECT) if parent else None
852+
853+
def render_with_model_dialect(node: exp.Expr, **overrides: t.Any) -> str:
854+
opts: t.Dict[str, t.Any] = {
855+
"dialect": meta_dialect,
856+
"pretty": self.pretty,
857+
"identify": self.identify,
858+
"normalize": self.normalize,
859+
"pad": self.pad,
860+
"indent": self._indent,
861+
"normalize_functions": self.normalize_functions,
862+
"leading_comma": self.leading_comma,
863+
"max_text_width": self.max_text_width,
864+
"comments": self.comments,
865+
}
866+
opts.update(overrides)
867+
868+
# Keep boolean literals anywhere in the value (audit args, physical_properties,
869+
# merge_filter, ...) as `TRUE`/`FALSE`: tsql would otherwise emit `(1 = 1)`,
870+
# which reformats differently on the next pass. The value is transpiled with
871+
# the model dialect anyway when it is used, e.g. in the rendered audit query.
872+
def keep_boolean_literal(n: exp.Expr) -> exp.Expr:
873+
if not isinstance(n, exp.Boolean):
874+
return n
875+
literal = exp.var("TRUE" if n.this else "FALSE")
876+
literal.comments = n.comments
877+
return literal
878+
879+
return node.transform(keep_boolean_literal).sql(**opts)
880+
742881
if isinstance(prop, MacroFunc):
743-
sql = self.indent(self.sql(prop, comment=False))
882+
# A macro in property position wraps user-authored arguments, so it carries
883+
# warehouse SQL the same way `columns` or `audits` do. Clear the outer node's
884+
# own comments (not `.this`'s, which `_macro_func_sql` already attaches)
885+
# before rendering with the model dialect, mirroring what `comment=False`
886+
# does for the non-dialect path below -- passing `comments=False` here
887+
# instead would build a fresh Generator with comments globally disabled,
888+
# silently dropping every comment in the subtree rather than just the
889+
# redundant outer one.
890+
if meta_dialect:
891+
prop_for_render = prop.copy()
892+
prop_for_render.comments = None
893+
sql = self.indent(render_with_model_dialect(prop_for_render))
894+
else:
895+
sql = self.indent(self.sql(prop, comment=False))
744896
else:
745-
sql = self.indent(f"{prop.name} {self.sql(prop, 'value')}")
897+
value = prop.args.get("value")
898+
899+
if (
900+
meta_dialect
901+
and isinstance(value, exp.Expr)
902+
and _meta_render_policy().get(prop.name.lower())
903+
):
904+
value_sql = render_with_model_dialect(value)
905+
elif (
906+
meta_dialect
907+
and isinstance(value, exp.Array)
908+
and _dialect_renders_array_as_brackets(meta_dialect)
909+
):
910+
# Dialect-agnostic properties (e.g. `tags`, `ignored_rules`) that hold a
911+
# list still go through the base (dialect=None) generator, which renders
912+
# an `exp.Array` as `ARRAY(...)`. On BigQuery `ARRAY(` is parsed as a
913+
# subquery constructor, so a multi-element `ARRAY('a', 'b')` fails to
914+
# reparse ("Required keyword: 'value' missing for Property"). Render it
915+
# as a bracketed list literal instead -- but only for dialects that
916+
# actually spell arrays that way; dialects that reuse `[`/`]` for
917+
# identifier quoting (tsql, sqlite, ...) keep the generic `ARRAY(...)`
918+
# form, which they parse back correctly. The elements themselves stay on
919+
# the dialect-agnostic path (`self.expressions`, not
920+
# `render_with_model_dialect`): these are SQLMesh's own scalar values
921+
# (tag/rule name strings), not user warehouse SQL, so they must not be
922+
# transpiled with the model dialect (e.g. tsql boolean literals turning
923+
# into `(1 = 1)`).
924+
value_sql = f"[{self.expressions(value, flat=True)}]"
925+
else:
926+
value_sql = self.sql(prop, "value")
927+
928+
sql = self.indent(f"{prop.name} {value_sql}")
746929

747930
if i < size - 1:
748931
sql += ","
@@ -853,11 +1036,29 @@ def format_model_expressions(
8531036
Returns:
8541037
A string representing the formatted model.
8551038
"""
1039+
1040+
def tag_meta_dialect(expression: exp.Expr) -> exp.Expr:
1041+
"""Record the model dialect on meta nodes so `_props_sql` can render the
1042+
warehouse-SQL properties (columns, audits, physical_properties, ...) with it
1043+
while the SQLMesh-owned ones stay dialect-agnostic. Tags nested ModelKind
1044+
nodes too, since kinds carry expression properties of their own such as
1045+
`time_data_type` and `unique_key`."""
1046+
if not dialect or not is_meta_expression(expression):
1047+
return expression
1048+
1049+
expression = expression.copy()
1050+
for node in expression.find_all(Model, Audit, Metric, ModelKind):
1051+
node.meta[_SQLMESH_META_DIALECT] = dialect
1052+
expression.meta[_SQLMESH_META_DIALECT] = dialect
1053+
return expression
1054+
8561055
if len(expressions) == 1 and is_meta_expression(expressions[0]):
8571056
# Meta expressions (MODEL/AUDIT/METRIC) are SQLMesh DDL, not standard SQL,
8581057
# so they must never be transpiled to the target dialect (e.g. tsql would
8591058
# rewrite a boolean property like `allow_partials TRUE` to `(1 = 1)`).
860-
return expressions[0].sql(
1059+
# Individual properties whose values *are* warehouse SQL still render with
1060+
# the model dialect -- see `_props_sql` / `_meta_render_policy`.
1061+
return tag_meta_dialect(expressions[0]).sql(
8611062
pretty=True, dialect=None, normalize_functions=normalize_functions
8621063
)
8631064

@@ -893,7 +1094,7 @@ def cast_to_colon(node: exp.Expr) -> exp.Expr:
8931094
return ";\n\n".join(
8941095
# Meta expressions (MODEL/AUDIT/METRIC) are SQLMesh DDL and must stay
8951096
# dialect-agnostic; only the actual query/statement expressions transpile.
896-
expression.sql(
1097+
tag_meta_dialect(expression).sql(
8971098
pretty=True,
8981099
dialect=None if is_meta_expression(expression) else dialect,
8991100
normalize_functions=normalize_functions,

0 commit comments

Comments
 (0)