mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-09-28 12:54:15 +02:00
Fix all remaining resource warnings, refs #693
https://gistpreview.github.io/?0bb8e869b82f6ff0db647de755182502
This commit is contained in:
parent
77359bea30
commit
dc9947a5e1
7 changed files with 110 additions and 61 deletions
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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():
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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():
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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):
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue