Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
81 changes: 38 additions & 43 deletions superset/sql/dialects/db2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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')."""
Expand Down Expand Up @@ -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

Expand Down
20 changes: 20 additions & 0 deletions tests/unit_tests/sql/dialects/db2_tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
Comment thread
aminghadersohi marked this conversation as resolved.
("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
Loading