diff --git a/superset/sql/dialects/db2.py b/superset/sql/dialects/db2.py index b10fd925567f..6211511c6b1e 100644 --- a/superset/sql/dialects/db2.py +++ b/superset/sql/dialects/db2.py @@ -28,6 +28,23 @@ from sqlglot.dialects.dialect import rename_func from sqlglot.dialects.postgres import Postgres +LABELED_DURATION_UNITS = { + "MICROSECOND", + "MICROSECONDS", + "SECOND", + "SECONDS", + "MINUTE", + "MINUTES", + "HOUR", + "HOURS", + "DAY", + "DAYS", + "MONTH", + "MONTHS", + "YEAR", + "YEARS", +} + class DB2Interval(exp.Expression): """DB2 labeled duration expression (e.g., '1 DAYS', '2 MONTHS').""" @@ -67,68 +84,46 @@ class Tokenizer(Postgres.Tokenizer): class Parser(Postgres.Parser): """DB2 SQL parser with support for labeled durations.""" - def _parse_term(self) -> exp.Expression | None: + def _parse_term(self, parse_mod: bool = True) -> exp.Expression | None: """ Override term parsing to support DB2 labeled durations. This is called during expression parsing for addition/subtraction operations. We intercept patterns like `expr + 1 DAYS` and parse them - specially. + specially. Everything else follows sqlglot's own implementation, + including the ``parse_mod`` flag it passes while parsing LIMIT and + OFFSET, and the other term operators (e.g. COLLATE). """ - this = self._parse_factor() - if not this: - return None - - while self._match_set((tokens.TokenType.PLUS, tokens.TokenType.DASH)): - op = self._prev.token_type + this = self._parse_factor(parse_mod=parse_mod) - # Parse the right side of the + or - - rhs = self._parse_factor() - if not rhs: # pragma: no cover - break + while self._match_set(self.TERM): + token_type = self._prev.token_type + klass = self.TERM[token_type] + comments = self._prev_comments + expression = self._parse_factor(parse_mod=parse_mod) # Check if there's a time unit after the right side # This handles patterns like: expr + 1 DAYS, expr + (func()) DAYS if ( - self._curr + token_type in (tokens.TokenType.PLUS, tokens.TokenType.DASH) + and expression is not None + and self._curr and self._curr.token_type == tokens.TokenType.VAR - and self._curr.text.upper() - in { - "MICROSECOND", - "MICROSECONDS", - "SECOND", - "SECONDS", - "MINUTE", - "MINUTES", - "HOUR", - "HOURS", - "DAY", - "DAYS", - "MONTH", - "MONTHS", - "YEAR", - "YEARS", - } + and self._curr.text.upper() in LABELED_DURATION_UNITS ): # Found a DB2 labeled duration unit_token = self._curr self._advance() - - duration = DB2Interval( - this=rhs, + expression = DB2Interval( + this=expression, unit=exp.Literal.string(unit_token.text.upper()), ) - if op == tokens.TokenType.PLUS: - this = exp.Add(this=this, expression=duration) - else: - this = exp.Sub(this=this, expression=duration) - else: - # Not a labeled duration - use normal Add/Sub - if op == tokens.TokenType.PLUS: - this = exp.Add(this=this, expression=rhs) - else: - this = exp.Sub(this=this, expression=rhs) + this = self.expression( + klass(this=this, expression=expression), comments=comments + ) + if isinstance(this, exp.Collate): + self._normalize_collate(this) return this diff --git a/tests/unit_tests/sql/dialects/db2_tests.py b/tests/unit_tests/sql/dialects/db2_tests.py index b8add1d65b1a..a6a4c7e97047 100644 --- a/tests/unit_tests/sql/dialects/db2_tests.py +++ b/tests/unit_tests/sql/dialects/db2_tests.py @@ -244,3 +244,23 @@ def test_column_plus_literal_duration() -> None: # Should parse as (col + 1 DAYS), not (col + 1) AS DAYS assert regenerated == "SELECT col + 1 DAYS FROM t" + + +@pytest.mark.parametrize( + ("sql", "expected"), + [ + ("SELECT * FROM t LIMIT 10 OFFSET 2", "SELECT * FROM t LIMIT 10 OFFSET 2"), + ( + "SELECT * FROM (SELECT 1 AS a FROM t) AS x LIMIT 5", + "SELECT * FROM (SELECT 1 AS a FROM t) AS x LIMIT 5", + ), + ("SELECT a % 3 FROM t LIMIT 4", "SELECT a % 3 FROM t LIMIT 4"), + ("SELECT * FROM t LIMIT 10 %", "SELECT * FROM t LIMIT 10 PERCENT"), + ("SELECT a COLLATE x FROM t", "SELECT a COLLATE x FROM t"), + ], +) +def test_limit_offset_and_other_term_operators(sql: str, expected: str) -> None: + """ + sqlglot passes ``parse_mod`` to ``_parse_term`` while parsing LIMIT/OFFSET. + """ + assert parse_one(sql, dialect=DB2).sql(dialect=DB2) == expected