Fix all remaining resource warnings, refs #693

https://gistpreview.github.io/?0bb8e869b82f6ff0db647de755182502
This commit is contained in:
Simon Willison 2025-12-11 16:46:05 -08:00
commit dc9947a5e1
7 changed files with 110 additions and 61 deletions

View file

@ -905,7 +905,7 @@ def insert_upsert_options(*, require_pk=False):
required=True, required=True,
), ),
click.argument("table"), click.argument("table"),
click.argument("file", type=click.File("rb"), required=True), click.argument("file", type=click.File("rb", lazy=True), required=True),
click.option( click.option(
"--pk", "--pk",
help="Columns to use as the primary key, e.g. id", help="Columns to use as the primary key, e.g. id",
@ -2000,6 +2000,7 @@ def memory(
for i, path in enumerate(paths): for i, path in enumerate(paths):
# Path may have a :format suffix # Path may have a :format suffix
fp = None fp = None
should_close_fp = False
if ":" in path and path.rsplit(":", 1)[-1].upper() in Format.__members__: if ":" in path and path.rsplit(":", 1)[-1].upper() in Format.__members__:
path, suffix = path.rsplit(":", 1) path, suffix = path.rsplit(":", 1)
format = Format[suffix.upper()] format = Format[suffix.upper()]
@ -2017,29 +2018,32 @@ def memory(
file_table = stem file_table = stem
stem_counts[stem] = stem_counts.get(stem, 1) + 1 stem_counts[stem] = stem_counts.get(stem, 1) + 1
fp = file_path.open("rb") fp = file_path.open("rb")
rows, format_used = rows_from_file(fp, format=format, encoding=encoding) should_close_fp = True
tracker = None try:
if format_used in (Format.CSV, Format.TSV) and not no_detect_types: rows, format_used = rows_from_file(fp, format=format, encoding=encoding)
tracker = TypeTracker() tracker = None
rows = tracker.wrap(rows) if format_used in (Format.CSV, Format.TSV) and not no_detect_types:
if flatten: tracker = TypeTracker()
rows = (_flatten(row) for row in rows) rows = tracker.wrap(rows)
if flatten:
rows = (_flatten(row) for row in rows)
db[file_table].insert_all(rows, alter=True) db[file_table].insert_all(rows, alter=True)
if tracker is not None: if tracker is not None:
db[file_table].transform(types=tracker.types) db[file_table].transform(types=tracker.types)
# Add convenient t / t1 / t2 views # Add convenient t / t1 / t2 views
view_names = ["t{}".format(i + 1)] view_names = ["t{}".format(i + 1)]
if i == 0: if i == 0:
view_names.append("t") view_names.append("t")
for view_name in view_names: for view_name in view_names:
if not db[view_name].exists(): if not db[view_name].exists():
db.create_view( db.create_view(
view_name, "select * from {}".format(quote_identifier(file_table)) view_name,
) "select * from {}".format(quote_identifier(file_table)),
)
if fp: finally:
fp.close() if should_close_fp and fp:
fp.close()
if analyze: if analyze:
_analyze(db, tables=None, columns=None, save=False) _analyze(db, tables=None, columns=None, save=False)

View file

@ -9,7 +9,29 @@ import json
import os import os
import sys import sys
from . import recipes from . import recipes
from typing import Dict, cast, BinaryIO, Iterable, Optional, Tuple, Type from typing import Dict, cast, BinaryIO, Iterable, Iterator, Optional, Tuple, Type
class _CloseableIterator(Iterator[dict]):
"""Iterator wrapper that closes a file when iteration is complete."""
def __init__(self, iterator: Iterator[dict], closeable: io.IOBase):
self._iterator = iterator
self._closeable = closeable
def __iter__(self) -> "_CloseableIterator":
return self
def __next__(self) -> dict:
try:
return next(self._iterator)
except StopIteration:
self._closeable.close()
raise
def close(self) -> None:
self._closeable.close()
import click import click
@ -299,7 +321,8 @@ def rows_from_file(
reader = csv.DictReader(decoded_fp, dialect=dialect) reader = csv.DictReader(decoded_fp, dialect=dialect)
else: else:
reader = csv.DictReader(decoded_fp) reader = csv.DictReader(decoded_fp)
return _extra_key_strategy(reader, ignore_extras, extras_key), Format.CSV rows = _extra_key_strategy(reader, ignore_extras, extras_key)
return _CloseableIterator(iter(rows), decoded_fp), Format.CSV
elif format == Format.TSV: elif format == Format.TSV:
rows = rows_from_file( rows = rows_from_file(
fp, format=Format.CSV, dialect=csv.excel_tab, encoding=encoding fp, format=Format.CSV, dialect=csv.excel_tab, encoding=encoding

View file

@ -19,9 +19,11 @@ def _supports_pragma_function_list():
db = Database(memory=True) db = Database(memory=True)
try: try:
db.execute("select * from pragma_function_list()") db.execute("select * from pragma_function_list()")
return True
except Exception: except Exception:
return False return False
return True finally:
db.close()
def _has_compiled_ext(): def _has_compiled_ext():

View file

@ -2,6 +2,14 @@ from sqlite_utils.db import Index, View, Database, XIndex, XIndexColumn
import pytest import pytest
def _check_supports_strict():
"""Check if SQLite supports strict tables without leaking the database."""
db = Database(memory=True)
result = db.supports_strict
db.close()
return result
def test_table_names(existing_db): def test_table_names(existing_db):
assert ["foo"] == existing_db.table_names() assert ["foo"] == existing_db.table_names()
@ -282,7 +290,7 @@ def test_use_rowid(fresh_db):
@pytest.mark.skipif( @pytest.mark.skipif(
not Database(memory=True).supports_strict, not _check_supports_strict(),
reason="Needs SQLite version that supports strict", reason="Needs SQLite version that supports strict",
) )
@pytest.mark.parametrize( @pytest.mark.parametrize(

View file

@ -9,9 +9,11 @@ def _supports_pragma_function_list():
db = Database(memory=True) db = Database(memory=True)
try: try:
db.execute("select * from pragma_function_list()") db.execute("select * from pragma_function_list()")
return True
except Exception: except Exception:
return False return False
return True finally:
db.close()
def test_register_commands(): def test_register_commands():

View file

@ -14,8 +14,11 @@ def test_recreate_ignored_for_in_memory():
def test_recreate_not_allowed_for_connection(): def test_recreate_not_allowed_for_connection():
conn = sqlite3.connect(":memory:") conn = sqlite3.connect(":memory:")
with pytest.raises(AssertionError): try:
Database(conn, recreate=True) with pytest.raises(AssertionError):
Database(conn, recreate=True)
finally:
conn.close()
@pytest.mark.parametrize( @pytest.mark.parametrize(

View file

@ -42,39 +42,46 @@ def test_register_function_deterministic(fresh_db):
def test_register_function_deterministic_tries_again_if_exception_raised(fresh_db): def test_register_function_deterministic_tries_again_if_exception_raised(fresh_db):
# Save the original connection so we can close it later
original_conn = fresh_db.conn
fresh_db.conn = MagicMock() fresh_db.conn = MagicMock()
fresh_db.conn.create_function = MagicMock() fresh_db.conn.create_function = MagicMock()
@fresh_db.register_function(deterministic=True) try:
def to_lower_2(s):
return s.lower()
fresh_db.conn.create_function.assert_called_with( @fresh_db.register_function(deterministic=True)
"to_lower_2", 1, to_lower_2, deterministic=True def to_lower_2(s):
) return s.lower()
first = True fresh_db.conn.create_function.assert_called_with(
"to_lower_2", 1, to_lower_2, deterministic=True
)
def side_effect(*args, **kwargs): first = True
# Raise exception only first time this is called
nonlocal first
if first:
first = False
raise sqlite3.NotSupportedError()
# But if sqlite3.NotSupportedError is raised, it tries again def side_effect(*args, **kwargs):
fresh_db.conn.create_function.reset_mock() # Raise exception only first time this is called
fresh_db.conn.create_function.side_effect = side_effect nonlocal first
if first:
first = False
raise sqlite3.NotSupportedError()
@fresh_db.register_function(deterministic=True) # But if sqlite3.NotSupportedError is raised, it tries again
def to_lower_3(s): fresh_db.conn.create_function.reset_mock()
return s.lower() fresh_db.conn.create_function.side_effect = side_effect
# Should have been called once with deterministic=True and once without @fresh_db.register_function(deterministic=True)
assert fresh_db.conn.create_function.call_args_list == [ def to_lower_3(s):
call("to_lower_3", 1, to_lower_3, deterministic=True), return s.lower()
call("to_lower_3", 1, to_lower_3),
] # Should have been called once with deterministic=True and once without
assert fresh_db.conn.create_function.call_args_list == [
call("to_lower_3", 1, to_lower_3, deterministic=True),
call("to_lower_3", 1, to_lower_3),
]
finally:
# Close the original connection that was replaced with the mock
original_conn.close()
def test_register_function_replace(fresh_db): def test_register_function_replace(fresh_db):