mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-09-17 22:14:09 +02:00
Fix type errors in db.py
- Add type annotation for Database.conn - Add type: ignore for optional sqlite_dump import - Update execute/query parameter types to Sequence|Dict for sqlite3 compatibility - Use getattr for fn.__name__ access to handle callables without __name__ - Handle None return from find_spatialite() with OSError - Fix pk_values assignment to use local variable 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
parent
5c15db90f4
commit
8a54c8b2d9
1 changed files with 16 additions and 11 deletions
|
|
@ -41,7 +41,7 @@ import uuid
|
||||||
from sqlite_utils.plugins import pm
|
from sqlite_utils.plugins import pm
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from sqlite_dump import iterdump
|
from sqlite_dump import iterdump # type: ignore[import-not-found]
|
||||||
except ImportError:
|
except ImportError:
|
||||||
iterdump = None
|
iterdump = None
|
||||||
|
|
||||||
|
|
@ -526,7 +526,7 @@ class Database:
|
||||||
self.execute(attach_sql)
|
self.execute(attach_sql)
|
||||||
|
|
||||||
def query(
|
def query(
|
||||||
self, sql: str, params: Optional[Union[Iterable, dict]] = None
|
self, sql: str, params: Optional[Union[Sequence, Dict[str, Any]]] = None
|
||||||
) -> Generator[dict, None, None]:
|
) -> Generator[dict, None, None]:
|
||||||
"""
|
"""
|
||||||
Execute ``sql`` and return an iterable of dictionaries representing each row.
|
Execute ``sql`` and return an iterable of dictionaries representing each row.
|
||||||
|
|
@ -541,7 +541,7 @@ class Database:
|
||||||
yield dict(zip(keys, row))
|
yield dict(zip(keys, row))
|
||||||
|
|
||||||
def execute(
|
def execute(
|
||||||
self, sql: str, parameters: Optional[Union[Iterable, dict]] = None
|
self, sql: str, parameters: Optional[Union[Sequence, Dict[str, Any]]] = None
|
||||||
) -> sqlite3.Cursor:
|
) -> sqlite3.Cursor:
|
||||||
"""
|
"""
|
||||||
Execute SQL query and return a ``sqlite3.Cursor``.
|
Execute SQL query and return a ``sqlite3.Cursor``.
|
||||||
|
|
@ -806,10 +806,11 @@ class Database:
|
||||||
:param tables: Subset list of tables to return counts for.
|
:param tables: Subset list of tables to return counts for.
|
||||||
"""
|
"""
|
||||||
sql = 'select "table", count from {}'.format(self._counts_table_name)
|
sql = 'select "table", count from {}'.format(self._counts_table_name)
|
||||||
if tables:
|
tables_list = list(tables) if tables else None
|
||||||
sql += ' where "table" in ({})'.format(", ".join("?" for table in tables))
|
if tables_list:
|
||||||
|
sql += ' where "table" in ({})'.format(", ".join("?" for _ in tables_list))
|
||||||
try:
|
try:
|
||||||
return {r[0]: r[1] for r in self.execute(sql, tables).fetchall()}
|
return {r[0]: r[1] for r in self.execute(sql, tables_list).fetchall()}
|
||||||
except OperationalError:
|
except OperationalError:
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|
@ -826,7 +827,7 @@ class Database:
|
||||||
)
|
)
|
||||||
|
|
||||||
def execute_returning_dicts(
|
def execute_returning_dicts(
|
||||||
self, sql: str, params: Optional[Union[Iterable, dict]] = None
|
self, sql: str, params: Optional[Union[Sequence, Dict[str, Any]]] = None
|
||||||
) -> List[dict]:
|
) -> List[dict]:
|
||||||
return list(self.query(sql, params))
|
return list(self.query(sql, params))
|
||||||
|
|
||||||
|
|
@ -1340,6 +1341,8 @@ class Database:
|
||||||
"""
|
"""
|
||||||
if path is None:
|
if path is None:
|
||||||
path = find_spatialite()
|
path = find_spatialite()
|
||||||
|
if path is None:
|
||||||
|
raise OSError("Could not find SpatiaLite extension")
|
||||||
|
|
||||||
self.conn.enable_load_extension(True)
|
self.conn.enable_load_extension(True)
|
||||||
self.conn.load_extension(path)
|
self.conn.load_extension(path)
|
||||||
|
|
@ -3006,7 +3009,7 @@ class Table(Queryable):
|
||||||
bar.update(1)
|
bar.update(1)
|
||||||
return jsonify_if_needed(fn(v))
|
return jsonify_if_needed(fn(v))
|
||||||
|
|
||||||
fn_name = fn.__name__
|
fn_name = getattr(fn, "__name__", "fn")
|
||||||
if fn_name == "<lambda>":
|
if fn_name == "<lambda>":
|
||||||
fn_name = f"lambda_{abs(hash(fn))}"
|
fn_name = f"lambda_{abs(hash(fn))}"
|
||||||
self.db.register_function(convert_value, name=fn_name)
|
self.db.register_function(convert_value, name=fn_name)
|
||||||
|
|
@ -3251,9 +3254,11 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
# We can populate .last_pk right here
|
# We can populate .last_pk right here
|
||||||
if num_records_processed == 1:
|
if num_records_processed == 1:
|
||||||
self.last_pk = tuple(record[pk] for pk in pks)
|
pk_values = tuple(record[pk] for pk in pks)
|
||||||
if len(self.last_pk) == 1:
|
if len(pk_values) == 1:
|
||||||
self.last_pk = self.last_pk[0]
|
self.last_pk = pk_values[0]
|
||||||
|
else:
|
||||||
|
self.last_pk = pk_values
|
||||||
return queries_and_params
|
return queries_and_params
|
||||||
|
|
||||||
def insert_chunk(
|
def insert_chunk(
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue