diff --git a/docs/source/api.rst b/docs/source/api.rst index 7c6fd70a..6c6c5c1b 100644 --- a/docs/source/api.rst +++ b/docs/source/api.rst @@ -38,6 +38,8 @@ The :meth:`~sqlparse.format` function accepts the following keyword arguments. ``truncate_strings`` If ``truncate_strings`` is a positive integer, string literals longer than the given value will be truncated. + Truncation does not split doubled quotes or backslash escape sequences; + the retained content may be shorter than the requested width. ``truncate_char`` (default: "[...]") If long string literals are truncated (see above) this value will be append diff --git a/sqlparse/filters/tokens.py b/sqlparse/filters/tokens.py index cc00a844..395ca4a4 100644 --- a/sqlparse/filters/tokens.py +++ b/sqlparse/filters/tokens.py @@ -5,6 +5,8 @@ # This module is part of python-sqlparse and is released under # the BSD License: https://opensource.org/licenses/BSD-3-Clause +import re + from sqlparse import tokens as T @@ -47,13 +49,14 @@ def process(self, stream): yield ttype, value continue - if value[:2] == "''": - inner = value[2:-2] - quote = "''" - else: - inner = value[1:-1] - quote = "'" - + inner = value[1:-1] if len(inner) > self.width: - value = ''.join((quote, inner[:self.width], self.char, quote)) + # Keep doubled quotes and backslash escapes intact at the + # truncation boundary, without exceeding the requested width. + end = 0 + for match in re.finditer(r"''|\\.|.", inner, re.DOTALL): + if match.end() > self.width: + break + end = match.end() + value = ''.join(("'", inner[:end], self.char, "'")) yield ttype, value diff --git a/tests/test_format.py b/tests/test_format.py index 93495067..27d6221f 100644 --- a/tests/test_format.py +++ b/tests/test_format.py @@ -738,6 +738,27 @@ def test_truncate_strings(): assert formatted == "update foo set value = 'xxxYYY';" +@pytest.mark.parametrize('marker', ['[...]', '...', '']) +@pytest.mark.parametrize(('literal', 'width', 'expected'), [ + ("'ab''cdef'", 3, "'ab[...]'"), + ("'ab''cdef'", 4, "'ab''[...]'"), + ("'''abcdef'", 3, "'''a[...]'"), + ("'''abcdef'", 6, "'''abcd[...]'"), + ("'a\nb''cdef'", 4, "'a\nb[...]'"), + ("'ab''''cdef'", 5, "'ab''[...]'"), + ("'ab''''cdef'", 6, "'ab''''[...]'"), + (r"'ab\'cdef'", 3, "'ab[...]'"), + (r"'ab\'cdef'", 4, r"'ab\'[...]'"), + (r"'ab\\cdef'", 3, "'ab[...]'"), + (r"'ab\\cdef'", 4, r"'ab\\[...]'"), + ("'ab''cd'", 6, "'ab''cd'"), +]) +def test_truncate_strings_preserves_escapes(literal, width, expected, marker): + formatted = sqlparse.format(f'SELECT {literal};', truncate_strings=width, + truncate_char=marker) + assert formatted == 'SELECT {};'.format(expected.replace('[...]', marker)) + + @pytest.mark.parametrize('option', ['bar', -1, 0]) def test_truncate_strings_invalid_option2(option): with pytest.raises(SQLParseError):