Skip to content
Open
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
14 changes: 9 additions & 5 deletions sqlparse/engine/grouping.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,8 +240,10 @@ def post(tlist, pidx, tidx, nidx):
return pidx, nidx

valid_prev = valid_next = valid
# Comments between an operator and operand belong to the comparison,
# not a separate identifier list (issue604).
_group(tlist, sql.Comparison, match,
valid_prev, valid_next, post, extend=False)
valid_prev, valid_next, post, extend=False, skip_cm=True)


@recurse(sql.Identifier)
Expand Down Expand Up @@ -481,7 +483,8 @@ def _group(tlist, cls, match,
post=None,
extend=True,
recurse=True,
depth=0
depth=0,
skip_cm=False
):
"""Groups together tokens that are joined by a middle token. i.e. x < y"""
if MAX_GROUPING_DEPTH is not None and depth > MAX_GROUPING_DEPTH:
Expand All @@ -505,15 +508,16 @@ def _group(tlist, cls, match,
if tidx < 0: # tidx shouldn't get negative
continue

if token.is_whitespace:
if token.is_whitespace or (skip_cm and imt(token, i=sql.Comment,
t=T.Comment)):
continue

if recurse and token.is_group and not isinstance(token, cls):
_group(token, cls, match, valid_prev, valid_next,
post, extend, True, depth + 1)
post, extend, True, depth + 1, skip_cm)

if match(token):
nidx, next_ = tlist.token_next(tidx)
nidx, next_ = tlist.token_next(tidx, skip_cm=skip_cm)
if prev_ and valid_prev(prev_) and valid_next(next_):
from_idx, to_idx = post(tlist, pidx, tidx, nidx)
grp = tlist.group_tokens(cls, from_idx, to_idx, extend=extend)
Expand Down
14 changes: 14 additions & 0 deletions tests/test_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,20 @@ def test_strip_comments_invalid_option(self):
with pytest.raises(SQLParseError):
sqlparse.format(sql, strip_comments=None)

@pytest.mark.parametrize('comment', (
'/* comment */', '/* first line\nsecond line */', '-- comment\n',
))
def test_reindent_commented_comparison_columns(self, comment):
# issue604: an ungrouped comparison split one SELECT list into
# nested lists, so every comment increased subsequent indentation.
text = f'SELECT leadcol = leadvalue, aaa = {comment} bbb, ccc = ccc, ddd = ddd FROM t'
result = sqlparse.format(text, strip_comments=True, reindent=True)
assert result == ('SELECT leadcol = leadvalue,\n'
' aaa = bbb,\n'
' ccc = ccc,\n'
' ddd = ddd\n'
'FROM t')

def test_strip_comments_multi(self):
sql = '/* sql starts here */\nselect'
res = sqlparse.format(sql, strip_comments=True)
Expand Down
49 changes: 49 additions & 0 deletions tests/test_grouping.py
Original file line number Diff line number Diff line change
Expand Up @@ -488,6 +488,55 @@ def test_comparison_with_keywords():
assert isinstance(p.tokens[0], sql.Comparison)


@pytest.mark.parametrize('left_comment,right_comment', (
('/* before */', ''),
('', '/* after */'),
('/* before */', '/* after */'),
('-- before\n', '-- after\n'),
('/* first */ /* second */', '/* third */'),
('/*+ hint */', '/*+ hint */'),
))
def test_comparison_with_comments(left_comment, right_comment):
text = f'foo {left_comment} = {right_comment} bar'
statement = sqlparse.parse(text)[0]
assert len(statement.tokens) == 1
comparison = statement.tokens[0]
assert isinstance(comparison, sql.Comparison)
assert comparison.left.get_real_name() == 'foo'
assert comparison.right.value == 'bar'
assert str(comparison) == text
assert any(token.ttype in T.Comment for token in comparison.flatten())


def test_comparison_with_comments_in_parenthesis():
statement = sqlparse.parse('(foo /* before */ = /* after */ bar)')[0]
comparison = statement.tokens[0].tokens[1]
assert isinstance(comparison, sql.Comparison)
assert comparison.left.get_real_name() == 'foo'
assert comparison.right.value == 'bar'


@pytest.mark.parametrize('text,right', (
('foo >= /* comment */ 25.5', '25.5'),
("foo NOT LIKE /* comment */ 'bar'", "'bar'"),
('foo = /* comment */ DATE(bar)', 'DATE(bar)'),
('foo = /* comment */ NULL', 'NULL'),
))
def test_comparison_with_commented_operands(text, right):
statement = sqlparse.parse(text)[0]
comparison = statement.tokens[0]
assert isinstance(comparison, sql.Comparison)
assert comparison.right.value == right
assert str(comparison) == text


def test_comparison_does_not_cross_statement_after_comment():
text = 'foo = /* comment */; SELECT bar'
first, second = sqlparse.parse(text)
assert not any(isinstance(token, sql.Comparison) for token in first.tokens)
assert str(first) + str(second) == text


def test_comparison_with_floats():
# issue145
p = sqlparse.parse('foo = 25.5')[0]
Expand Down