From bb68f25cd63ab76c2bb532ab18e0e49f7457d5e5 Mon Sep 17 00:00:00 2001 From: adenzhou1350 <209601943+adenzhou1350@users.noreply.github.com> Date: Tue, 6 Oct 2026 16:17:27 +0800 Subject: [PATCH] Accept binary file-like SQL streams --- sqlparse/__init__.py | 4 ++-- sqlparse/lexer.py | 4 ++-- tests/test_binary_stream.py | 47 +++++++++++++++++++++++++++++++++++++ 3 files changed, 51 insertions(+), 4 deletions(-) create mode 100644 tests/test_binary_stream.py diff --git a/sqlparse/__init__.py b/sqlparse/__init__.py index 8411f5a3..6744a0af 100644 --- a/sqlparse/__init__.py +++ b/sqlparse/__init__.py @@ -30,11 +30,11 @@ def parse( def parsestream( - stream: str | IO[str], encoding: str | None = None + stream: str | bytes | IO[str] | IO[bytes], encoding: str | None = None ) -> Generator[sql.Statement, None, None]: """Parses sql statements from file-like object. - :param stream: A file-like object. + :param stream: SQL text or a text or binary file-like object. :param encoding: The encoding of the stream contents (optional). :returns: A generator of :class:`~sqlparse.sql.Statement` instances. """ diff --git a/sqlparse/lexer.py b/sqlparse/lexer.py index 5d4a5f97..0e569905 100644 --- a/sqlparse/lexer.py +++ b/sqlparse/lexer.py @@ -12,7 +12,7 @@ # http://pygments.org/ # It's separated from the rest of pygments to increase performance # and to allow some customizations. -from io import TextIOBase +from io import IOBase from threading import Lock from sqlparse import keywords, tokens @@ -116,7 +116,7 @@ def get_tokens(self, text, encoding=None): ``stack`` is the initial stack (default: ``['root']``) """ - if isinstance(text, TextIOBase): + if isinstance(text, IOBase): text = text.read() if isinstance(text, str): diff --git a/tests/test_binary_stream.py b/tests/test_binary_stream.py new file mode 100644 index 00000000..b738a19e --- /dev/null +++ b/tests/test_binary_stream.py @@ -0,0 +1,47 @@ +"""Text and binary streams should share the existing decoding path.""" +import gzip +from io import BytesIO, StringIO + +import pytest + +import sqlparse + + +@pytest.mark.parametrize('encoding', ['utf-8', 'latin-1']) +@pytest.mark.parametrize('compressed', [False, True]) +def test_parsestream_binary(encoding, compressed): + # issue352: GzipFile is a binary stream, not a TextIOBase. + text = "SELECT 'caf\N{LATIN SMALL LETTER E WITH ACUTE}'; SELECT 2;" + data = text.encode(encoding) + if compressed: + stream = gzip.GzipFile(fileobj=BytesIO(gzip.compress(data))) + else: + stream = BytesIO(data) + with stream: + result = [str(stmt) for stmt in sqlparse.parsestream(stream, encoding)] + assert not stream.closed + assert result == [str(stmt) for stmt in sqlparse.parse(text)] + + +def test_parsestream_binary_default_encoding(): + text = "SELECT '\N{CJK UNIFIED IDEOGRAPH-65E5}';" + with BytesIO(text.encode('utf-8')) as stream: + assert [str(stmt) for stmt in sqlparse.parsestream(stream)] == [text] + + +def test_parsestream_text(): + text = "SELECT 'caf\N{LATIN SMALL LETTER E WITH ACUTE}';" + with StringIO(text) as stream: + assert [str(stmt) for stmt in sqlparse.parsestream(stream)] == [text] + assert not stream.closed + + +def test_parsestream_binary_bad_encoding(): + with BytesIO(b"SELECT '\xff';") as stream: + with pytest.raises(UnicodeDecodeError): + list(sqlparse.parsestream(stream, encoding='utf-8')) + + +def test_parsestream_invalid_object(): + with pytest.raises(TypeError, match='Expected text or file-like object'): + list(sqlparse.parsestream(object()))