diff --git a/sqlmesh/core/dialect.py b/sqlmesh/core/dialect.py index 67918b6d14..f96887a61c 100644 --- a/sqlmesh/core/dialect.py +++ b/sqlmesh/core/dialect.py @@ -782,6 +782,10 @@ def _parse_interval_span(self: Parser, this: exp.Expr) -> exp.Interval: def _override(klass: t.Type[Tokenizer | Parser], func: t.Callable) -> None: name = func.__name__ + if getattr(klass, name, None) is func: + # Already overridden. Re-applying would save the override itself as the + # "original", making the wrapper call itself and recurse infinitely. + return setattr(klass, f"_{name}", getattr(klass, name)) setattr(klass, name, func) @@ -1170,11 +1174,12 @@ def extend_sqlglot() -> None: MacroDef, ) - generator.UNWRAPPED_INTERVAL_VALUES = ( - *generator.UNWRAPPED_INTERVAL_VALUES, - MacroStrReplace, - MacroVar, - ) + if MacroVar not in generator.UNWRAPPED_INTERVAL_VALUES: + generator.UNWRAPPED_INTERVAL_VALUES = ( + *generator.UNWRAPPED_INTERVAL_VALUES, + MacroStrReplace, + MacroVar, + ) _override(Parser, _parse_select) _override(Parser, _parse_statement) diff --git a/tests/core/test_dialect.py b/tests/core/test_dialect.py index e2f1daba3d..6cbd9107a0 100644 --- a/tests/core/test_dialect.py +++ b/tests/core/test_dialect.py @@ -1185,3 +1185,17 @@ def test_pipe_syntax(): ast.sql("bigquery") == "SELECT * FROM (WITH __tmp1 AS (SELECT id FROM t2) SELECT * FROM __tmp1)" ) + + +def test_extend_sqlglot_is_idempotent(): + # extend_sqlglot() runs at import time; calling it again must not re-wrap the + # already-installed overrides, otherwise they call themselves (RecursionError). + from sqlglot.generator import Generator + + before = Generator.UNWRAPPED_INTERVAL_VALUES + + d.extend_sqlglot() + + assert parse_one("SELECT CAST(1 AS INT)").sql() == "SELECT CAST(1 AS INT)" + # The class-level registries must not grow on repeated calls. + assert Generator.UNWRAPPED_INTERVAL_VALUES == before