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
4 changes: 2 additions & 2 deletions sqlparse/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
"""
Expand Down
4 changes: 2 additions & 2 deletions sqlparse/lexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down
47 changes: 47 additions & 0 deletions tests/test_binary_stream.py
Original file line number Diff line number Diff line change
@@ -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()))