mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-09-28 04:44:26 +02:00
Fix resource leaks caused by buffered readers
Spotted running 'just test -Wall' on the LLM project.
This commit is contained in:
parent
85b1be10c8
commit
6bc1d33d58
2 changed files with 75 additions and 1 deletions
|
|
@ -380,10 +380,14 @@ def rows_from_file(
|
||||||
"rows_from_file() requires a file-like object that supports peek(), such as io.BytesIO"
|
"rows_from_file() requires a file-like object that supports peek(), such as io.BytesIO"
|
||||||
)
|
)
|
||||||
if not first_bytes:
|
if not first_bytes:
|
||||||
|
buffered.close()
|
||||||
return (), Format.CSV
|
return (), Format.CSV
|
||||||
if first_bytes.startswith((b"[", b"{")):
|
if first_bytes.startswith((b"[", b"{")):
|
||||||
|
# JSON is read eagerly, so the detection wrapper can close now,
|
||||||
|
# including when parsing raises an error.
|
||||||
# TODO: Detect newline-JSON
|
# TODO: Detect newline-JSON
|
||||||
return rows_from_file(buffered, format=Format.JSON)
|
with buffered:
|
||||||
|
return rows_from_file(buffered, format=Format.JSON)
|
||||||
else:
|
else:
|
||||||
dialect = csv.Sniffer().sniff(
|
dialect = csv.Sniffer().sniff(
|
||||||
first_bytes.decode(encoding or "utf-8-sig", "ignore")
|
first_bytes.decode(encoding or "utf-8-sig", "ignore")
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,5 @@
|
||||||
|
import io
|
||||||
|
import json
|
||||||
from io import BytesIO, StringIO
|
from io import BytesIO, StringIO
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
@ -61,3 +63,71 @@ def test_rows_from_file_error_on_string_io():
|
||||||
assert ex.value.args == (
|
assert ex.value.args == (
|
||||||
"rows_from_file() requires a file-like object that supports peek(), such as io.BytesIO",
|
"rows_from_file() requires a file-like object that supports peek(), such as io.BytesIO",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def buffered_readers(monkeypatch):
|
||||||
|
# Keep wrappers alive so these checks cannot pass due to garbage collection.
|
||||||
|
readers = []
|
||||||
|
original = io.BufferedReader
|
||||||
|
|
||||||
|
def buffered_reader(*args, **kwargs):
|
||||||
|
reader = original(*args, **kwargs)
|
||||||
|
readers.append(reader)
|
||||||
|
return reader
|
||||||
|
|
||||||
|
monkeypatch.setattr(io, "BufferedReader", buffered_reader)
|
||||||
|
yield readers
|
||||||
|
for reader in readers:
|
||||||
|
reader.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"content, expected_format, expected_rows",
|
||||||
|
[
|
||||||
|
(b'[{"id": 1}]', Format.JSON, [{"id": 1}]),
|
||||||
|
(b'{"id": 1}', Format.JSON, [{"id": 1}]),
|
||||||
|
(b"[]", Format.JSON, []),
|
||||||
|
(b"", Format.CSV, []),
|
||||||
|
(b" \n\t", Format.CSV, []),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_detect_format_closes_eager_reader(
|
||||||
|
tmp_path, buffered_readers, content, expected_format, expected_rows
|
||||||
|
):
|
||||||
|
path = tmp_path / "input"
|
||||||
|
path.write_bytes(content)
|
||||||
|
with path.open("rb") as fp:
|
||||||
|
rows, detected = rows_from_file(fp)
|
||||||
|
assert detected == expected_format
|
||||||
|
assert list(rows) == expected_rows
|
||||||
|
assert len(buffered_readers) == 1
|
||||||
|
assert buffered_readers[0].closed
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("content", [b"[", b'{"id":'])
|
||||||
|
def test_detect_format_closes_reader_on_invalid_json(
|
||||||
|
tmp_path, buffered_readers, content
|
||||||
|
):
|
||||||
|
path = tmp_path / "input.json"
|
||||||
|
path.write_bytes(content)
|
||||||
|
with path.open("rb") as fp:
|
||||||
|
with pytest.raises(json.JSONDecodeError):
|
||||||
|
rows_from_file(fp)
|
||||||
|
assert len(buffered_readers) == 1
|
||||||
|
assert buffered_readers[0].closed
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"delimiter, expected_format", [(b",", Format.CSV), (b"\t", Format.TSV)]
|
||||||
|
)
|
||||||
|
def test_detect_format_keeps_streaming_reader_open(
|
||||||
|
buffered_readers, delimiter, expected_format
|
||||||
|
):
|
||||||
|
content = b"id" + delimiter + b"name\n" + (b"1" + delimiter + b"Cleo\n") * 2000
|
||||||
|
rows, detected = rows_from_file(BytesIO(content))
|
||||||
|
assert detected == expected_format
|
||||||
|
assert len(buffered_readers) == 1
|
||||||
|
assert not buffered_readers[0].closed
|
||||||
|
assert len(list(rows)) == 2000
|
||||||
|
assert buffered_readers[0].closed
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue