diff --git a/CHANGELOG b/CHANGELOG index 44d5938e..4dbaada1 100644 --- a/CHANGELOG +++ b/CHANGELOG @@ -1,7 +1,7 @@ Development Version ------------------- -Nothing yet. +* Fix CASE...END split-level tracking outside BEGIN blocks (pr839). Release 0.6.0 (Aug 13, 2026) diff --git a/sqlparse/engine/statement_splitter.py b/sqlparse/engine/statement_splitter.py index bc57d170..f691edcf 100644 --- a/sqlparse/engine/statement_splitter.py +++ b/sqlparse/engine/statement_splitter.py @@ -145,6 +145,10 @@ def _change_splitlevel(self, ttype, value): if res is not None: return res + if unified == 'CASE': + self._block_stack.append('CASE') + return 1 + # Handle closing keywords return self._handle_closing_keyword(unified) diff --git a/tests/test_split.py b/tests/test_split.py index 92c3fefe..fec09e91 100644 --- a/tests/test_split.py +++ b/tests/test_split.py @@ -374,3 +374,44 @@ def test_split_standalone_for_update(): assert stmts[1] == "SELECT 3;" +def test_splitlevel_case_end(): + # CASE in a plain SELECT did not increment the level, but its matching END + # decremented it unconditionally. This led to levels being wrong after the + # CASE WHEN ... END block. + s = sqlparse.engine.statement_splitter.StatementSplitter() + level = s.level + + token_stream = [ + (sqlparse.tokens.Keyword.DML, 'SELECT'), + (sqlparse.tokens.Keyword, 'CASE'), + (sqlparse.tokens.Keyword, 'WHEN'), + (sqlparse.tokens.Name, 'foo'), + (sqlparse.tokens.Keyword, 'THEN'), + (sqlparse.tokens.Number, '1'), + (sqlparse.tokens.Keyword, 'END'), + (sqlparse.tokens.Keyword, 'FROM'), + (sqlparse.tokens.Name, 't'), + ] + + for ttype, value in token_stream: + level += s._change_splitlevel(ttype, value) + + assert level == 0 + + # This issue could lead to incorrectly treating a semicolon inside a text + # literal as a statement terminator and incorrectly splitting the query. + assert len(sqlparse.parse( + "SELECT CASE WHEN 1 THEN 2 END, test IN ('foo \\', 'foo;') FROM t" + )) == 1 + + +@pytest.mark.parametrize('expression', [ + 'CASE WHEN 1 THEN 2 END', + 'CASE WHEN 1 THEN CASE WHEN 2 THEN 3 END ELSE 4 END', +]) +def test_split_case_end_escaped_semicolon(expression): + query = "SELECT " + expression + ", test IN ('foo \\', 'foo;') FROM t" + statements = sqlparse.parse(query + '; SELECT 2;') + assert [str(statement) for statement in statements] == [ + query + '; ', 'SELECT 2;', + ]