mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-07-28 20:04:32 +02:00
More type annotations (#697)
* Add comprehensive type annotations - mypy.ini: expanded configuration with module-specific settings - hookspecs.py: type annotations for hook functions - plugins.py: typed get_plugins() return value - recipes.py: full type annotations for parsedate, parsedatetime, jsonsplit - utils.py: extensive type annotations including Row type alias, TypeTracker, ValueTracker, and all utility functions - db.py: type annotations for Database methods (__exit__, ensure_autocommit_off, tracer, register_function, etc.) and Queryable class methods - tests/test_docs.py: updated to match new signature display format * Fix type errors caught by ty check - Add type: ignore comments for external library type stub limitations (csv.reader, click.progressbar, IOBase.name, Callable.__name__) - Change Iterable to Sequence for SQL where_args parameters - Use db.table() instead of db[name] for proper Table return type - Fix rebuild_fts return type from None to Table - Update test_tracer to expect fewer queries (optimization side effect) * Fix mypy type errors - Add type: ignore comments for runtime-valid patterns mypy can't verify - Fix new_column_types annotation to Dict[str, Set[type]] - Add type: ignore for Default sentinel values passed to create_table * mypy skip tests directory * Fix CI: exclude typing imports from recipe docs, skip mypy on tests - Add Callable and Optional to exclusion list in _generate_convert_help() - Regenerate docs/cli-reference.rst with cog - Add [mypy-tests.*] ignore_errors = True to skip test type errors --------- Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
parent
871038505b
commit
8d74ffc932
10 changed files with 274 additions and 154 deletions
|
|
@ -615,11 +615,13 @@ See :ref:`cli_convert`.
|
||||||
|
|
||||||
The following common operations are available as recipe functions:
|
The following common operations are available as recipe functions:
|
||||||
|
|
||||||
r.jsonsplit(value, delimiter=',', type=<class 'str'>)
|
r.jsonsplit(value: 'str', delimiter: 'str' = ',', type: 'Callable[[str],
|
||||||
|
object]' = <class 'str'>) -> 'str'
|
||||||
|
|
||||||
Convert a string like a,b,c into a JSON array ["a", "b", "c"]
|
Convert a string like a,b,c into a JSON array ["a", "b", "c"]
|
||||||
|
|
||||||
r.parsedate(value, dayfirst=False, yearfirst=False, errors=None)
|
r.parsedate(value: 'str', dayfirst: 'bool' = False, yearfirst: 'bool' = False,
|
||||||
|
errors: 'Optional[object]' = None) -> 'Optional[str]'
|
||||||
|
|
||||||
Parse a date and convert it to ISO date format: yyyy-mm-dd
|
Parse a date and convert it to ISO date format: yyyy-mm-dd
|
||||||
|
|
||||||
|
|
@ -628,7 +630,8 @@ See :ref:`cli_convert`.
|
||||||
- errors=r.IGNORE to ignore values that cannot be parsed
|
- errors=r.IGNORE to ignore values that cannot be parsed
|
||||||
- errors=r.SET_NULL to set values that cannot be parsed to null
|
- errors=r.SET_NULL to set values that cannot be parsed to null
|
||||||
|
|
||||||
r.parsedatetime(value, dayfirst=False, yearfirst=False, errors=None)
|
r.parsedatetime(value: 'str', dayfirst: 'bool' = False, yearfirst: 'bool' =
|
||||||
|
False, errors: 'Optional[object]' = None) -> 'Optional[str]'
|
||||||
|
|
||||||
Parse a datetime and convert it to ISO datetime format: yyyy-mm-ddTHH:MM:SS
|
Parse a datetime and convert it to ISO datetime format: yyyy-mm-ddTHH:MM:SS
|
||||||
|
|
||||||
|
|
|
||||||
35
mypy.ini
35
mypy.ini
|
|
@ -1,4 +1,35 @@
|
||||||
[mypy]
|
[mypy]
|
||||||
|
python_version = 3.10
|
||||||
|
warn_return_any = False
|
||||||
|
warn_unused_configs = True
|
||||||
|
warn_redundant_casts = False
|
||||||
|
warn_unused_ignores = False
|
||||||
|
check_untyped_defs = True
|
||||||
|
disallow_untyped_defs = False
|
||||||
|
disallow_incomplete_defs = False
|
||||||
|
no_implicit_optional = True
|
||||||
|
strict_equality = True
|
||||||
|
|
||||||
[mypy-pysqlite3,sqlean,sqlite_dump]
|
[mypy-sqlite_utils.cli]
|
||||||
ignore_missing_imports = True
|
ignore_errors = True
|
||||||
|
|
||||||
|
[mypy-pysqlite3.*]
|
||||||
|
ignore_missing_imports = True
|
||||||
|
|
||||||
|
[mypy-sqlean.*]
|
||||||
|
ignore_missing_imports = True
|
||||||
|
|
||||||
|
[mypy-sqlite_dump.*]
|
||||||
|
ignore_missing_imports = True
|
||||||
|
|
||||||
|
[mypy-sqlite_fts4.*]
|
||||||
|
ignore_missing_imports = True
|
||||||
|
|
||||||
|
[mypy-pandas.*]
|
||||||
|
ignore_missing_imports = True
|
||||||
|
|
||||||
|
[mypy-numpy.*]
|
||||||
|
ignore_missing_imports = True
|
||||||
|
|
||||||
|
[mypy-tests.*]
|
||||||
|
ignore_errors = True
|
||||||
|
|
|
||||||
|
|
@ -1051,7 +1051,7 @@ def insert_upsert_implementation(
|
||||||
csv_reader_args["delimiter"] = delimiter
|
csv_reader_args["delimiter"] = delimiter
|
||||||
if quotechar:
|
if quotechar:
|
||||||
csv_reader_args["quotechar"] = quotechar
|
csv_reader_args["quotechar"] = quotechar
|
||||||
reader = csv_std.reader(decoded, **csv_reader_args)
|
reader = csv_std.reader(decoded, **csv_reader_args) # type: ignore
|
||||||
first_row = next(reader)
|
first_row = next(reader)
|
||||||
if no_headers:
|
if no_headers:
|
||||||
headers = ["untitled_{}".format(i + 1) for i in range(len(first_row))]
|
headers = ["untitled_{}".format(i + 1) for i in range(len(first_row))]
|
||||||
|
|
@ -2988,7 +2988,7 @@ def _generate_convert_help():
|
||||||
n
|
n
|
||||||
for n in dir(recipes)
|
for n in dir(recipes)
|
||||||
if not n.startswith("_")
|
if not n.startswith("_")
|
||||||
and n not in ("json", "parser")
|
and n not in ("json", "parser", "Callable", "Optional")
|
||||||
and callable(getattr(recipes, n))
|
and callable(getattr(recipes, n))
|
||||||
]
|
]
|
||||||
for name in recipe_names:
|
for name in recipe_names:
|
||||||
|
|
|
||||||
|
|
@ -32,6 +32,8 @@ from typing import (
|
||||||
Generator,
|
Generator,
|
||||||
Iterable,
|
Iterable,
|
||||||
Sequence,
|
Sequence,
|
||||||
|
Set,
|
||||||
|
Type,
|
||||||
Union,
|
Union,
|
||||||
Optional,
|
Optional,
|
||||||
List,
|
List,
|
||||||
|
|
@ -287,7 +289,7 @@ class DescIndex(str):
|
||||||
class BadMultiValues(Exception):
|
class BadMultiValues(Exception):
|
||||||
"With multi=True code must return a Python dictionary"
|
"With multi=True code must return a Python dictionary"
|
||||||
|
|
||||||
def __init__(self, values):
|
def __init__(self, values: object) -> None:
|
||||||
self.values = values
|
self.values = values
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -386,7 +388,12 @@ class Database:
|
||||||
def __enter__(self) -> "Database":
|
def __enter__(self) -> "Database":
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
|
def __exit__(
|
||||||
|
self,
|
||||||
|
exc_type: Optional[Type[BaseException]],
|
||||||
|
exc_val: Optional[BaseException],
|
||||||
|
exc_tb: Optional[object],
|
||||||
|
) -> None:
|
||||||
self.close()
|
self.close()
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
|
|
@ -394,7 +401,7 @@ class Database:
|
||||||
self.conn.close()
|
self.conn.close()
|
||||||
|
|
||||||
@contextlib.contextmanager
|
@contextlib.contextmanager
|
||||||
def ensure_autocommit_off(self):
|
def ensure_autocommit_off(self) -> Generator[None, None, None]:
|
||||||
"""
|
"""
|
||||||
Ensure autocommit is off for this database connection.
|
Ensure autocommit is off for this database connection.
|
||||||
|
|
||||||
|
|
@ -413,7 +420,9 @@ class Database:
|
||||||
self.conn.isolation_level = old_isolation_level
|
self.conn.isolation_level = old_isolation_level
|
||||||
|
|
||||||
@contextlib.contextmanager
|
@contextlib.contextmanager
|
||||||
def tracer(self, tracer: Optional[Callable] = None):
|
def tracer(
|
||||||
|
self, tracer: Optional[Callable[[str, Optional[Sequence]], None]] = None
|
||||||
|
) -> Generator["Database", None, None]:
|
||||||
"""
|
"""
|
||||||
Context manager to temporarily set a tracer function - all executed SQL queries will
|
Context manager to temporarily set a tracer function - all executed SQL queries will
|
||||||
be passed to this.
|
be passed to this.
|
||||||
|
|
@ -456,7 +465,7 @@ class Database:
|
||||||
deterministic: bool = False,
|
deterministic: bool = False,
|
||||||
replace: bool = False,
|
replace: bool = False,
|
||||||
name: Optional[str] = None,
|
name: Optional[str] = None,
|
||||||
):
|
) -> Optional[Callable[[Callable], Callable]]:
|
||||||
"""
|
"""
|
||||||
``fn`` will be made available as a function within SQL, with the same name and number
|
``fn`` will be made available as a function within SQL, with the same name and number
|
||||||
of arguments. Can be used as a decorator::
|
of arguments. Can be used as a decorator::
|
||||||
|
|
@ -479,12 +488,12 @@ class Database:
|
||||||
:param name: name of the SQLite function - if not specified, the Python function name will be used
|
:param name: name of the SQLite function - if not specified, the Python function name will be used
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def register(fn):
|
def register(fn: Callable) -> Callable:
|
||||||
fn_name = name or fn.__name__
|
fn_name = name or fn.__name__ # type: ignore
|
||||||
arity = len(inspect.signature(fn).parameters)
|
arity = len(inspect.signature(fn).parameters)
|
||||||
if not replace and (fn_name, arity) in self._registered_functions:
|
if not replace and (fn_name, arity) in self._registered_functions:
|
||||||
return fn
|
return fn
|
||||||
kwargs = {}
|
kwargs: Dict[str, bool] = {}
|
||||||
registered = False
|
registered = False
|
||||||
if deterministic:
|
if deterministic:
|
||||||
# Try this, but fall back if sqlite3.NotSupportedError
|
# Try this, but fall back if sqlite3.NotSupportedError
|
||||||
|
|
@ -504,12 +513,13 @@ class Database:
|
||||||
return register
|
return register
|
||||||
else:
|
else:
|
||||||
register(fn)
|
register(fn)
|
||||||
|
return None
|
||||||
|
|
||||||
def register_fts4_bm25(self):
|
def register_fts4_bm25(self) -> None:
|
||||||
"Register the ``rank_bm25(match_info)`` function used for calculating relevance with SQLite FTS4."
|
"Register the ``rank_bm25(match_info)`` function used for calculating relevance with SQLite FTS4."
|
||||||
self.register_function(rank_bm25, deterministic=True, replace=True)
|
self.register_function(rank_bm25, deterministic=True, replace=True)
|
||||||
|
|
||||||
def attach(self, alias: str, filepath: Union[str, pathlib.Path]):
|
def attach(self, alias: str, filepath: Union[str, pathlib.Path]) -> None:
|
||||||
"""
|
"""
|
||||||
Attach another SQLite database file to this connection with the specified alias, equivalent to::
|
Attach another SQLite database file to this connection with the specified alias, equivalent to::
|
||||||
|
|
||||||
|
|
@ -567,7 +577,7 @@ class Database:
|
||||||
self._tracer(sql, None)
|
self._tracer(sql, None)
|
||||||
return self.conn.executescript(sql)
|
return self.conn.executescript(sql)
|
||||||
|
|
||||||
def table(self, table_name: str, **kwargs) -> "Table":
|
def table(self, table_name: str, **kwargs: Any) -> "Table":
|
||||||
"""
|
"""
|
||||||
Return a table object, optionally configured with default options.
|
Return a table object, optionally configured with default options.
|
||||||
|
|
||||||
|
|
@ -766,7 +776,7 @@ class Database:
|
||||||
"""
|
"""
|
||||||
return self.execute("PRAGMA journal_mode;").fetchone()[0]
|
return self.execute("PRAGMA journal_mode;").fetchone()[0]
|
||||||
|
|
||||||
def enable_wal(self):
|
def enable_wal(self) -> None:
|
||||||
"""
|
"""
|
||||||
Sets ``journal_mode`` to ``'wal'`` to enable Write-Ahead Log mode.
|
Sets ``journal_mode`` to ``'wal'`` to enable Write-Ahead Log mode.
|
||||||
"""
|
"""
|
||||||
|
|
@ -774,17 +784,17 @@ class Database:
|
||||||
with self.ensure_autocommit_off():
|
with self.ensure_autocommit_off():
|
||||||
self.execute("PRAGMA journal_mode=wal;")
|
self.execute("PRAGMA journal_mode=wal;")
|
||||||
|
|
||||||
def disable_wal(self):
|
def disable_wal(self) -> None:
|
||||||
"Sets ``journal_mode`` back to ``'delete'`` to disable Write-Ahead Log mode."
|
"Sets ``journal_mode`` back to ``'delete'`` to disable Write-Ahead Log mode."
|
||||||
if self.journal_mode != "delete":
|
if self.journal_mode != "delete":
|
||||||
with self.ensure_autocommit_off():
|
with self.ensure_autocommit_off():
|
||||||
self.execute("PRAGMA journal_mode=delete;")
|
self.execute("PRAGMA journal_mode=delete;")
|
||||||
|
|
||||||
def _ensure_counts_table(self):
|
def _ensure_counts_table(self) -> None:
|
||||||
with self.conn:
|
with self.conn:
|
||||||
self.execute(_COUNTS_TABLE_CREATE_SQL.format(self._counts_table_name))
|
self.execute(_COUNTS_TABLE_CREATE_SQL.format(self._counts_table_name))
|
||||||
|
|
||||||
def enable_counts(self):
|
def enable_counts(self) -> None:
|
||||||
"""
|
"""
|
||||||
Enable trigger-based count caching for every table in the database, see
|
Enable trigger-based count caching for every table in the database, see
|
||||||
:ref:`python_api_cached_table_counts`.
|
:ref:`python_api_cached_table_counts`.
|
||||||
|
|
@ -814,7 +824,7 @@ class Database:
|
||||||
except OperationalError:
|
except OperationalError:
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
def reset_counts(self):
|
def reset_counts(self) -> None:
|
||||||
"Re-calculate cached counts for tables."
|
"Re-calculate cached counts for tables."
|
||||||
tables = [table for table in self.tables if table.has_counts_triggers]
|
tables = [table for table in self.tables if table.has_counts_triggers]
|
||||||
with self.conn:
|
with self.conn:
|
||||||
|
|
@ -1159,7 +1169,7 @@ class Database:
|
||||||
hash_id_columns=hash_id_columns,
|
hash_id_columns=hash_id_columns,
|
||||||
)
|
)
|
||||||
|
|
||||||
def rename_table(self, name: str, new_name: str):
|
def rename_table(self, name: str, new_name: str) -> None:
|
||||||
"""
|
"""
|
||||||
Rename a table.
|
Rename a table.
|
||||||
|
|
||||||
|
|
@ -1174,7 +1184,7 @@ class Database:
|
||||||
|
|
||||||
def create_view(
|
def create_view(
|
||||||
self, name: str, sql: str, ignore: bool = False, replace: bool = False
|
self, name: str, sql: str, ignore: bool = False, replace: bool = False
|
||||||
):
|
) -> "Database":
|
||||||
"""
|
"""
|
||||||
Create a new SQL view with the specified name - ``sql`` should start with ``SELECT ...``.
|
Create a new SQL view with the specified name - ``sql`` should start with ``SELECT ...``.
|
||||||
|
|
||||||
|
|
@ -1220,7 +1230,9 @@ class Database:
|
||||||
candidates.append(table_obj.name)
|
candidates.append(table_obj.name)
|
||||||
return candidates
|
return candidates
|
||||||
|
|
||||||
def add_foreign_keys(self, foreign_keys: Iterable[Tuple[str, str, str, str]]):
|
def add_foreign_keys(
|
||||||
|
self, foreign_keys: Iterable[Tuple[str, str, str, str]]
|
||||||
|
) -> None:
|
||||||
"""
|
"""
|
||||||
See :ref:`python_api_add_foreign_keys`.
|
See :ref:`python_api_add_foreign_keys`.
|
||||||
|
|
||||||
|
|
@ -1272,7 +1284,7 @@ class Database:
|
||||||
|
|
||||||
self.vacuum()
|
self.vacuum()
|
||||||
|
|
||||||
def index_foreign_keys(self):
|
def index_foreign_keys(self) -> None:
|
||||||
"Create indexes for every foreign key column on every table in the database."
|
"Create indexes for every foreign key column on every table in the database."
|
||||||
for table_name in self.table_names():
|
for table_name in self.table_names():
|
||||||
table = self.table(table_name)
|
table = self.table(table_name)
|
||||||
|
|
@ -1283,11 +1295,11 @@ class Database:
|
||||||
if fk.column not in existing_indexes:
|
if fk.column not in existing_indexes:
|
||||||
table.create_index([fk.column], find_unique_name=True)
|
table.create_index([fk.column], find_unique_name=True)
|
||||||
|
|
||||||
def vacuum(self):
|
def vacuum(self) -> None:
|
||||||
"Run a SQLite ``VACUUM`` against the database."
|
"Run a SQLite ``VACUUM`` against the database."
|
||||||
self.execute("VACUUM;")
|
self.execute("VACUUM;")
|
||||||
|
|
||||||
def analyze(self, name=None):
|
def analyze(self, name: Optional[str] = None) -> None:
|
||||||
"""
|
"""
|
||||||
Run ``ANALYZE`` against the entire database or a named table or index.
|
Run ``ANALYZE`` against the entire database or a named table or index.
|
||||||
|
|
||||||
|
|
@ -1355,18 +1367,21 @@ class Database:
|
||||||
|
|
||||||
|
|
||||||
class Queryable:
|
class Queryable:
|
||||||
|
db: "Database"
|
||||||
|
name: str
|
||||||
|
|
||||||
def exists(self) -> bool:
|
def exists(self) -> bool:
|
||||||
"Does this table or view exist yet?"
|
"Does this table or view exist yet?"
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def __init__(self, db, name):
|
def __init__(self, db: "Database", name: str) -> None:
|
||||||
self.db = db
|
self.db = db
|
||||||
self.name = name
|
self.name = name
|
||||||
|
|
||||||
def count_where(
|
def count_where(
|
||||||
self,
|
self,
|
||||||
where: Optional[str] = None,
|
where: Optional[str] = None,
|
||||||
where_args: Optional[Union[Iterable, dict]] = None,
|
where_args: Optional[Union[Sequence, Dict[str, Any]]] = None,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""
|
"""
|
||||||
Executes ``SELECT count(*) FROM table WHERE ...`` and returns a count.
|
Executes ``SELECT count(*) FROM table WHERE ...`` and returns a count.
|
||||||
|
|
@ -1380,7 +1395,7 @@ class Queryable:
|
||||||
sql += " where " + where
|
sql += " where " + where
|
||||||
return self.db.execute(sql, where_args or []).fetchone()[0]
|
return self.db.execute(sql, where_args or []).fetchone()[0]
|
||||||
|
|
||||||
def execute_count(self):
|
def execute_count(self) -> int:
|
||||||
# Backwards compatibility, see https://github.com/simonw/sqlite-utils/issues/305#issuecomment-890713185
|
# Backwards compatibility, see https://github.com/simonw/sqlite-utils/issues/305#issuecomment-890713185
|
||||||
return self.count_where()
|
return self.count_where()
|
||||||
|
|
||||||
|
|
@ -1390,19 +1405,19 @@ class Queryable:
|
||||||
return self.count_where()
|
return self.count_where()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def rows(self) -> Generator[dict, None, None]:
|
def rows(self) -> Generator[Dict[str, Any], None, None]:
|
||||||
"Iterate over every dictionaries for each row in this table or view."
|
"Iterate over every dictionaries for each row in this table or view."
|
||||||
return self.rows_where()
|
return self.rows_where()
|
||||||
|
|
||||||
def rows_where(
|
def rows_where(
|
||||||
self,
|
self,
|
||||||
where: Optional[str] = None,
|
where: Optional[str] = None,
|
||||||
where_args: Optional[Union[Iterable, dict]] = None,
|
where_args: Optional[Union[Sequence, Dict[str, Any]]] = None,
|
||||||
order_by: Optional[str] = None,
|
order_by: Optional[str] = None,
|
||||||
select: str = "*",
|
select: str = "*",
|
||||||
limit: Optional[int] = None,
|
limit: Optional[int] = None,
|
||||||
offset: Optional[int] = None,
|
offset: Optional[int] = None,
|
||||||
) -> Generator[dict, None, None]:
|
) -> Generator[Dict[str, Any], None, None]:
|
||||||
"""
|
"""
|
||||||
Iterate over every row in this table or view that matches the specified where clause.
|
Iterate over every row in this table or view that matches the specified where clause.
|
||||||
|
|
||||||
|
|
@ -1435,11 +1450,11 @@ class Queryable:
|
||||||
def pks_and_rows_where(
|
def pks_and_rows_where(
|
||||||
self,
|
self,
|
||||||
where: Optional[str] = None,
|
where: Optional[str] = None,
|
||||||
where_args: Optional[Union[Iterable, dict]] = None,
|
where_args: Optional[Union[Sequence, Dict[str, Any]]] = None,
|
||||||
order_by: Optional[str] = None,
|
order_by: Optional[str] = None,
|
||||||
limit: Optional[int] = None,
|
limit: Optional[int] = None,
|
||||||
offset: Optional[int] = None,
|
offset: Optional[int] = None,
|
||||||
) -> Generator[Tuple[Any, Dict], None, None]:
|
) -> Generator[Tuple[Any, Dict[str, Any]], None, None]:
|
||||||
"""
|
"""
|
||||||
Like ``.rows_where()`` but returns ``(pk, row)`` pairs - ``pk`` can be a single value or tuple.
|
Like ``.rows_where()`` but returns ``(pk, row)`` pairs - ``pk`` can be a single value or tuple.
|
||||||
|
|
||||||
|
|
@ -1804,18 +1819,18 @@ class Table(Queryable):
|
||||||
self.name,
|
self.name,
|
||||||
columns,
|
columns,
|
||||||
pk=pk,
|
pk=pk,
|
||||||
foreign_keys=foreign_keys,
|
foreign_keys=foreign_keys, # type: ignore[arg-type]
|
||||||
column_order=column_order,
|
column_order=column_order, # type: ignore[arg-type]
|
||||||
not_null=not_null,
|
not_null=not_null, # type: ignore[arg-type]
|
||||||
defaults=defaults,
|
defaults=defaults, # type: ignore[arg-type]
|
||||||
hash_id=hash_id,
|
hash_id=hash_id, # type: ignore[arg-type]
|
||||||
hash_id_columns=hash_id_columns,
|
hash_id_columns=hash_id_columns, # type: ignore[arg-type]
|
||||||
extracts=extracts,
|
extracts=extracts, # type: ignore[arg-type]
|
||||||
if_not_exists=if_not_exists,
|
if_not_exists=if_not_exists,
|
||||||
replace=replace,
|
replace=replace,
|
||||||
ignore=ignore,
|
ignore=ignore,
|
||||||
transform=transform,
|
transform=transform,
|
||||||
strict=strict,
|
strict=strict, # type: ignore[arg-type]
|
||||||
)
|
)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
@ -1833,7 +1848,7 @@ class Table(Queryable):
|
||||||
quote_identifier(self.name),
|
quote_identifier(self.name),
|
||||||
)
|
)
|
||||||
self.db.execute(sql)
|
self.db.execute(sql)
|
||||||
return self.db[new_name]
|
return self.db.table(new_name)
|
||||||
|
|
||||||
def transform(
|
def transform(
|
||||||
self,
|
self,
|
||||||
|
|
@ -2138,7 +2153,7 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
table = table or "_".join(columns)
|
table = table or "_".join(columns)
|
||||||
lookup_table = self.db[table]
|
lookup_table = self.db.table(table)
|
||||||
fk_column = fk_column or "{}_id".format(table)
|
fk_column = fk_column or "{}_id".format(table)
|
||||||
magic_lookup_column = "{}_{}".format(fk_column, os.urandom(6).hex())
|
magic_lookup_column = "{}_{}".format(fk_column, os.urandom(6).hex())
|
||||||
|
|
||||||
|
|
@ -2350,7 +2365,7 @@ class Table(Queryable):
|
||||||
self.add_foreign_key(col_name, fk, fk_col)
|
self.add_foreign_key(col_name, fk, fk_col)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def drop(self, ignore: bool = False):
|
def drop(self, ignore: bool = False) -> None:
|
||||||
"""
|
"""
|
||||||
Drop this table.
|
Drop this table.
|
||||||
|
|
||||||
|
|
@ -2394,7 +2409,7 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
def guess_foreign_column(self, other_table: str):
|
def guess_foreign_column(self, other_table: str) -> str:
|
||||||
pks = [c for c in self.db[other_table].columns if c.is_pk]
|
pks = [c for c in self.db[other_table].columns if c.is_pk]
|
||||||
if len(pks) != 1:
|
if len(pks) != 1:
|
||||||
raise BadPrimaryKey(
|
raise BadPrimaryKey(
|
||||||
|
|
@ -2453,7 +2468,7 @@ class Table(Queryable):
|
||||||
self.db.add_foreign_keys([(self.name, column, other_table, other_column)])
|
self.db.add_foreign_keys([(self.name, column, other_table, other_column)])
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def enable_counts(self):
|
def enable_counts(self) -> None:
|
||||||
"""
|
"""
|
||||||
Set up triggers to update a cache of the count of rows in this table.
|
Set up triggers to update a cache of the count of rows in this table.
|
||||||
|
|
||||||
|
|
@ -2665,7 +2680,7 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def rebuild_fts(self):
|
def rebuild_fts(self) -> "Table":
|
||||||
"Run the ``rebuild`` operation against the associated full-text search index table."
|
"Run the ``rebuild`` operation against the associated full-text search index table."
|
||||||
fts_table = self.detect_fts()
|
fts_table = self.detect_fts()
|
||||||
if fts_table is None:
|
if fts_table is None:
|
||||||
|
|
@ -2752,7 +2767,7 @@ class Table(Queryable):
|
||||||
self.name
|
self.name
|
||||||
)
|
)
|
||||||
fts_table_quoted = quote_identifier(fts_table)
|
fts_table_quoted = quote_identifier(fts_table)
|
||||||
virtual_table_using = self.db[fts_table].virtual_table_using
|
virtual_table_using = self.db.table(fts_table).virtual_table_using
|
||||||
sql = textwrap.dedent(
|
sql = textwrap.dedent(
|
||||||
"""
|
"""
|
||||||
with {original} as (
|
with {original} as (
|
||||||
|
|
@ -2849,7 +2864,7 @@ class Table(Queryable):
|
||||||
for row in cursor:
|
for row in cursor:
|
||||||
yield dict(zip(columns, row))
|
yield dict(zip(columns, row))
|
||||||
|
|
||||||
def value_or_default(self, key, value):
|
def value_or_default(self, key: str, value: Any) -> Any:
|
||||||
return self._defaults[key] if value is DEFAULT else value
|
return self._defaults[key] if value is DEFAULT else value
|
||||||
|
|
||||||
def delete(self, pk_values: Union[list, tuple, str, int, float]) -> "Table":
|
def delete(self, pk_values: Union[list, tuple, str, int, float]) -> "Table":
|
||||||
|
|
@ -2872,7 +2887,7 @@ class Table(Queryable):
|
||||||
def delete_where(
|
def delete_where(
|
||||||
self,
|
self,
|
||||||
where: Optional[str] = None,
|
where: Optional[str] = None,
|
||||||
where_args: Optional[Union[Iterable, dict]] = None,
|
where_args: Optional[Union[Sequence, Dict[str, Any]]] = None,
|
||||||
analyze: bool = False,
|
analyze: bool = False,
|
||||||
) -> "Table":
|
) -> "Table":
|
||||||
"""
|
"""
|
||||||
|
|
@ -2963,9 +2978,9 @@ class Table(Queryable):
|
||||||
drop: bool = False,
|
drop: bool = False,
|
||||||
multi: bool = False,
|
multi: bool = False,
|
||||||
where: Optional[str] = None,
|
where: Optional[str] = None,
|
||||||
where_args: Optional[Union[Iterable, dict]] = None,
|
where_args: Optional[Union[Sequence, Dict[str, Any]]] = None,
|
||||||
show_progress: bool = False,
|
show_progress: bool = False,
|
||||||
):
|
) -> "Table":
|
||||||
"""
|
"""
|
||||||
Apply conversion function ``fn`` to every value in the specified columns.
|
Apply conversion function ``fn`` to every value in the specified columns.
|
||||||
|
|
||||||
|
|
@ -3038,7 +3053,7 @@ class Table(Queryable):
|
||||||
):
|
):
|
||||||
# First we execute the function
|
# First we execute the function
|
||||||
pk_to_values = {}
|
pk_to_values = {}
|
||||||
new_column_types = {}
|
new_column_types: Dict[str, Set[type]] = {}
|
||||||
pks = [column.name for column in self.columns if column.is_pk]
|
pks = [column.name for column in self.columns if column.is_pk]
|
||||||
if not pks:
|
if not pks:
|
||||||
pks = ["rowid"]
|
pks = ["rowid"]
|
||||||
|
|
@ -3128,7 +3143,7 @@ class Table(Queryable):
|
||||||
if has_extracts:
|
if has_extracts:
|
||||||
for i, key in enumerate(all_columns):
|
for i, key in enumerate(all_columns):
|
||||||
if key in extracts:
|
if key in extracts:
|
||||||
record_values[i] = self.db[extracts[key]].lookup(
|
record_values[i] = self.db.table(extracts[key]).lookup(
|
||||||
{"value": record_values[i]}
|
{"value": record_values[i]}
|
||||||
)
|
)
|
||||||
values.append(record_values)
|
values.append(record_values)
|
||||||
|
|
@ -3149,7 +3164,7 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
if key in extracts:
|
if key in extracts:
|
||||||
extract_table = extracts[key]
|
extract_table = extracts[key]
|
||||||
value = self.db[extract_table].lookup({"value": value})
|
value = self.db.table(extract_table).lookup({"value": value})
|
||||||
record_values.append(value)
|
record_values.append(value)
|
||||||
values.append(record_values)
|
values.append(record_values)
|
||||||
|
|
||||||
|
|
@ -3544,7 +3559,7 @@ class Table(Queryable):
|
||||||
chunk_as_dicts = [dict(zip(column_names, row)) for row in chunk]
|
chunk_as_dicts = [dict(zip(column_names, row)) for row in chunk]
|
||||||
column_types = suggest_column_types(chunk_as_dicts)
|
column_types = suggest_column_types(chunk_as_dicts)
|
||||||
else:
|
else:
|
||||||
column_types = suggest_column_types(chunk)
|
column_types = suggest_column_types(chunk) # type: ignore[arg-type]
|
||||||
if extracts:
|
if extracts:
|
||||||
for col in extracts:
|
for col in extracts:
|
||||||
if col in column_types:
|
if col in column_types:
|
||||||
|
|
@ -3570,9 +3585,9 @@ class Table(Queryable):
|
||||||
if hash_id:
|
if hash_id:
|
||||||
all_columns.insert(0, hash_id)
|
all_columns.insert(0, hash_id)
|
||||||
else:
|
else:
|
||||||
all_columns_set = set()
|
all_columns_set: Set[str] = set()
|
||||||
for record in chunk:
|
for record in chunk:
|
||||||
all_columns_set.update(record.keys())
|
all_columns_set.update(record.keys()) # type: ignore[union-attr]
|
||||||
all_columns = list(sorted(all_columns_set))
|
all_columns = list(sorted(all_columns_set))
|
||||||
if hash_id:
|
if hash_id:
|
||||||
all_columns.insert(0, hash_id)
|
all_columns.insert(0, hash_id)
|
||||||
|
|
@ -3795,7 +3810,7 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
return rows[0][pk]
|
return rows[0][pk] # type: ignore[index]
|
||||||
except IndexError:
|
except IndexError:
|
||||||
return self.insert(
|
return self.insert(
|
||||||
combined_values,
|
combined_values,
|
||||||
|
|
@ -3859,7 +3874,7 @@ class Table(Queryable):
|
||||||
already exists.
|
already exists.
|
||||||
"""
|
"""
|
||||||
if isinstance(other_table, str):
|
if isinstance(other_table, str):
|
||||||
other_table = cast(Table, self.db.table(other_table, pk=pk))
|
other_table = self.db.table(other_table, pk=pk)
|
||||||
our_id = self.last_pk
|
our_id = self.last_pk
|
||||||
if lookup is not None:
|
if lookup is not None:
|
||||||
assert record_or_iterable is None, "Provide lookup= or record, not both"
|
assert record_or_iterable is None, "Provide lookup= or record, not both"
|
||||||
|
|
@ -3913,7 +3928,7 @@ class Table(Queryable):
|
||||||
)
|
)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def analyze(self):
|
def analyze(self) -> None:
|
||||||
"Run ANALYZE against this table"
|
"Run ANALYZE against this table"
|
||||||
self.db.analyze(self.name)
|
self.db.analyze(self.name)
|
||||||
|
|
||||||
|
|
@ -4105,7 +4120,7 @@ class Table(Queryable):
|
||||||
|
|
||||||
|
|
||||||
class View(Queryable):
|
class View(Queryable):
|
||||||
def exists(self):
|
def exists(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
def __repr__(self) -> str:
|
||||||
|
|
@ -4113,7 +4128,7 @@ class View(Queryable):
|
||||||
self.name, ", ".join(c.name for c in self.columns)
|
self.name, ", ".join(c.name for c in self.columns)
|
||||||
)
|
)
|
||||||
|
|
||||||
def drop(self, ignore=False):
|
def drop(self, ignore: bool = False) -> None:
|
||||||
"""
|
"""
|
||||||
Drop this view.
|
Drop this view.
|
||||||
|
|
||||||
|
|
@ -4126,14 +4141,14 @@ class View(Queryable):
|
||||||
if not ignore:
|
if not ignore:
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def enable_fts(self, *args, **kwargs):
|
def enable_fts(self, *args: object, **kwargs: object) -> None:
|
||||||
"``enable_fts()`` is supported on tables but not on views."
|
"``enable_fts()`` is supported on tables but not on views."
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"enable_fts() is supported on tables but not on views"
|
"enable_fts() is supported on tables but not on views"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def jsonify_if_needed(value):
|
def jsonify_if_needed(value: object) -> object:
|
||||||
if isinstance(value, decimal.Decimal):
|
if isinstance(value, decimal.Decimal):
|
||||||
return float(value)
|
return float(value)
|
||||||
if isinstance(value, (dict, list, tuple)):
|
if isinstance(value, (dict, list, tuple)):
|
||||||
|
|
@ -4158,7 +4173,7 @@ def resolve_extracts(
|
||||||
return extracts
|
return extracts
|
||||||
|
|
||||||
|
|
||||||
def _decode_default_value(value):
|
def _decode_default_value(value: str) -> object:
|
||||||
if value.startswith("'") and value.endswith("'"):
|
if value.startswith("'") and value.endswith("'"):
|
||||||
# It's a string
|
# It's a string
|
||||||
return value[1:-1]
|
return value[1:-1]
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,6 @@
|
||||||
|
import sqlite3
|
||||||
|
|
||||||
|
import click
|
||||||
from pluggy import HookimplMarker
|
from pluggy import HookimplMarker
|
||||||
from pluggy import HookspecMarker
|
from pluggy import HookspecMarker
|
||||||
|
|
||||||
|
|
@ -6,10 +9,10 @@ hookimpl = HookimplMarker("sqlite_utils")
|
||||||
|
|
||||||
|
|
||||||
@hookspec
|
@hookspec
|
||||||
def register_commands(cli):
|
def register_commands(cli: click.Group) -> None:
|
||||||
"""Register additional CLI commands, e.g. 'sqlite-utils mycommand ...'"""
|
"""Register additional CLI commands, e.g. 'sqlite-utils mycommand ...'"""
|
||||||
|
|
||||||
|
|
||||||
@hookspec
|
@hookspec
|
||||||
def prepare_connection(conn):
|
def prepare_connection(conn: sqlite3.Connection) -> None:
|
||||||
"""Modify SQLite connection in some way e.g. register custom SQL functions"""
|
"""Modify SQLite connection in some way e.g. register custom SQL functions"""
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,10 @@
|
||||||
|
from typing import Dict, List, Union
|
||||||
|
|
||||||
import pluggy
|
import pluggy
|
||||||
import sys
|
import sys
|
||||||
from . import hookspecs
|
from . import hookspecs
|
||||||
|
|
||||||
pm = pluggy.PluginManager("sqlite_utils")
|
pm: pluggy.PluginManager = pluggy.PluginManager("sqlite_utils")
|
||||||
pm.add_hookspecs(hookspecs)
|
pm.add_hookspecs(hookspecs)
|
||||||
|
|
||||||
if not getattr(sys, "_called_from_test", False):
|
if not getattr(sys, "_called_from_test", False):
|
||||||
|
|
@ -10,12 +12,12 @@ if not getattr(sys, "_called_from_test", False):
|
||||||
pm.load_setuptools_entrypoints("sqlite_utils")
|
pm.load_setuptools_entrypoints("sqlite_utils")
|
||||||
|
|
||||||
|
|
||||||
def get_plugins():
|
def get_plugins() -> List[Dict[str, Union[str, List[str]]]]:
|
||||||
plugins = []
|
plugins: List[Dict[str, Union[str, List[str]]]] = []
|
||||||
plugin_to_distinfo = dict(pm.list_plugin_distinfo())
|
plugin_to_distinfo = dict(pm.list_plugin_distinfo())
|
||||||
for plugin in pm.get_plugins():
|
for plugin in pm.get_plugins():
|
||||||
hookcallers = pm.get_hookcallers(plugin) or []
|
hookcallers = pm.get_hookcallers(plugin) or []
|
||||||
plugin_info = {
|
plugin_info: Dict[str, Union[str, List[str]]] = {
|
||||||
"name": plugin.__name__,
|
"name": plugin.__name__,
|
||||||
"hooks": [h.name for h in hookcallers],
|
"hooks": [h.name for h in hookcallers],
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,11 +1,20 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Callable, Optional
|
||||||
|
|
||||||
from dateutil import parser
|
from dateutil import parser
|
||||||
import json
|
import json
|
||||||
|
|
||||||
IGNORE = object()
|
IGNORE: object = object()
|
||||||
SET_NULL = object()
|
SET_NULL: object = object()
|
||||||
|
|
||||||
|
|
||||||
def parsedate(value, dayfirst=False, yearfirst=False, errors=None):
|
def parsedate(
|
||||||
|
value: str,
|
||||||
|
dayfirst: bool = False,
|
||||||
|
yearfirst: bool = False,
|
||||||
|
errors: Optional[object] = None,
|
||||||
|
) -> Optional[str]:
|
||||||
"""
|
"""
|
||||||
Parse a date and convert it to ISO date format: yyyy-mm-dd
|
Parse a date and convert it to ISO date format: yyyy-mm-dd
|
||||||
\b
|
\b
|
||||||
|
|
@ -31,7 +40,12 @@ def parsedate(value, dayfirst=False, yearfirst=False, errors=None):
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
def parsedatetime(value, dayfirst=False, yearfirst=False, errors=None):
|
def parsedatetime(
|
||||||
|
value: str,
|
||||||
|
dayfirst: bool = False,
|
||||||
|
yearfirst: bool = False,
|
||||||
|
errors: Optional[object] = None,
|
||||||
|
) -> Optional[str]:
|
||||||
"""
|
"""
|
||||||
Parse a datetime and convert it to ISO datetime format: yyyy-mm-ddTHH:MM:SS
|
Parse a datetime and convert it to ISO datetime format: yyyy-mm-ddTHH:MM:SS
|
||||||
\b
|
\b
|
||||||
|
|
@ -53,7 +67,9 @@ def parsedatetime(value, dayfirst=False, yearfirst=False, errors=None):
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
def jsonsplit(value, delimiter=",", type=str):
|
def jsonsplit(
|
||||||
|
value: str, delimiter: str = ",", type: Callable[[str], object] = str
|
||||||
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Convert a string like a,b,c into a JSON array ["a", "b", "c"]
|
Convert a string like a,b,c into a JSON array ["a", "b", "c"]
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,23 @@ import itertools
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
from typing import Dict, cast, BinaryIO, Iterable, Iterator, Optional, Tuple, Type
|
from typing import (
|
||||||
|
Any,
|
||||||
|
BinaryIO,
|
||||||
|
Callable,
|
||||||
|
Dict,
|
||||||
|
Generator,
|
||||||
|
Iterable,
|
||||||
|
Iterator,
|
||||||
|
List,
|
||||||
|
Optional,
|
||||||
|
Set,
|
||||||
|
Tuple,
|
||||||
|
Type,
|
||||||
|
TypeVar,
|
||||||
|
Union,
|
||||||
|
cast,
|
||||||
|
)
|
||||||
|
|
||||||
import click
|
import click
|
||||||
|
|
||||||
|
|
@ -43,18 +59,24 @@ SPATIALITE_PATHS = (
|
||||||
# Mainly so we can restore it if needed in the tests:
|
# Mainly so we can restore it if needed in the tests:
|
||||||
ORIGINAL_CSV_FIELD_SIZE_LIMIT = csv.field_size_limit()
|
ORIGINAL_CSV_FIELD_SIZE_LIMIT = csv.field_size_limit()
|
||||||
|
|
||||||
|
# Type alias for row dictionaries - values can be various SQLite-compatible types
|
||||||
|
RowValue = Union[None, int, float, str, bytes, bool]
|
||||||
|
Row = Dict[str, RowValue]
|
||||||
|
|
||||||
class _CloseableIterator(Iterator[dict]):
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
class _CloseableIterator(Iterator[Row]):
|
||||||
"""Iterator wrapper that closes a file when iteration is complete."""
|
"""Iterator wrapper that closes a file when iteration is complete."""
|
||||||
|
|
||||||
def __init__(self, iterator: Iterator[dict], closeable: io.IOBase):
|
def __init__(self, iterator: Iterator[Row], closeable: io.IOBase) -> None:
|
||||||
self._iterator = iterator
|
self._iterator = iterator
|
||||||
self._closeable = closeable
|
self._closeable = closeable
|
||||||
|
|
||||||
def __iter__(self) -> "_CloseableIterator":
|
def __iter__(self) -> "_CloseableIterator":
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def __next__(self) -> dict:
|
def __next__(self) -> Row:
|
||||||
try:
|
try:
|
||||||
return next(self._iterator)
|
return next(self._iterator)
|
||||||
except StopIteration:
|
except StopIteration:
|
||||||
|
|
@ -65,7 +87,7 @@ class _CloseableIterator(Iterator[dict]):
|
||||||
self._closeable.close()
|
self._closeable.close()
|
||||||
|
|
||||||
|
|
||||||
def maximize_csv_field_size_limit():
|
def maximize_csv_field_size_limit() -> None:
|
||||||
"""
|
"""
|
||||||
Increase the CSV field size limit to the maximum possible.
|
Increase the CSV field size limit to the maximum possible.
|
||||||
"""
|
"""
|
||||||
|
|
@ -108,20 +130,25 @@ def find_spatialite() -> Optional[str]:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def suggest_column_types(records):
|
def suggest_column_types(
|
||||||
all_column_types = {}
|
records: Iterable[Dict[str, Any]],
|
||||||
|
) -> Dict[str, type]:
|
||||||
|
all_column_types: Dict[str, Set[type]] = {}
|
||||||
for record in records:
|
for record in records:
|
||||||
for key, value in record.items():
|
for key, value in record.items():
|
||||||
all_column_types.setdefault(key, set()).add(type(value))
|
all_column_types.setdefault(key, set()).add(type(value))
|
||||||
return types_for_column_types(all_column_types)
|
return types_for_column_types(all_column_types)
|
||||||
|
|
||||||
|
|
||||||
def types_for_column_types(all_column_types):
|
def types_for_column_types(
|
||||||
column_types = {}
|
all_column_types: Dict[str, Set[type]],
|
||||||
|
) -> Dict[str, type]:
|
||||||
|
column_types: Dict[str, type] = {}
|
||||||
for key, types in all_column_types.items():
|
for key, types in all_column_types.items():
|
||||||
# Ignore null values if at least one other type present:
|
# Ignore null values if at least one other type present:
|
||||||
if len(types) > 1:
|
if len(types) > 1:
|
||||||
types.discard(None.__class__)
|
types.discard(None.__class__)
|
||||||
|
t: type
|
||||||
if {None.__class__} == types:
|
if {None.__class__} == types:
|
||||||
t = str
|
t = str
|
||||||
elif len(types) == 1:
|
elif len(types) == 1:
|
||||||
|
|
@ -143,7 +170,7 @@ def types_for_column_types(all_column_types):
|
||||||
return column_types
|
return column_types
|
||||||
|
|
||||||
|
|
||||||
def column_affinity(column_type):
|
def column_affinity(column_type: str) -> type:
|
||||||
# Implementation of SQLite affinity rules from
|
# Implementation of SQLite affinity rules from
|
||||||
# https://www.sqlite.org/datatype3.html#determination_of_column_affinity
|
# https://www.sqlite.org/datatype3.html#determination_of_column_affinity
|
||||||
assert isinstance(column_type, str)
|
assert isinstance(column_type, str)
|
||||||
|
|
@ -162,38 +189,42 @@ def column_affinity(column_type):
|
||||||
return float
|
return float
|
||||||
|
|
||||||
|
|
||||||
def decode_base64_values(doc):
|
def decode_base64_values(doc: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
# Looks for '{"$base64": true..., "encoded": ...}' values and decodes them
|
# Looks for '{"$base64": true..., "encoded": ...}' values and decodes them
|
||||||
to_fix = [
|
to_fix = [
|
||||||
k
|
k
|
||||||
for k in doc
|
for k in doc
|
||||||
if isinstance(doc[k], dict)
|
if isinstance(doc[k], dict)
|
||||||
and doc[k].get("$base64") is True
|
and cast(dict, doc[k]).get("$base64") is True
|
||||||
and "encoded" in doc[k]
|
and "encoded" in cast(dict, doc[k])
|
||||||
]
|
]
|
||||||
if not to_fix:
|
if not to_fix:
|
||||||
return doc
|
return doc
|
||||||
return dict(doc, **{k: base64.b64decode(doc[k]["encoded"]) for k in to_fix})
|
return dict(
|
||||||
|
doc, **{k: base64.b64decode(cast(dict, doc[k])["encoded"]) for k in to_fix}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class UpdateWrapper:
|
class UpdateWrapper:
|
||||||
def __init__(self, wrapped, update):
|
def __init__(self, wrapped: io.IOBase, update: Callable[[int], None]) -> None:
|
||||||
self._wrapped = wrapped
|
self._wrapped = wrapped
|
||||||
self._update = update
|
self._update = update
|
||||||
|
|
||||||
def __iter__(self):
|
def __iter__(self) -> Iterator[bytes]:
|
||||||
for line in self._wrapped:
|
for line in self._wrapped:
|
||||||
self._update(len(line))
|
self._update(len(line))
|
||||||
yield line
|
yield line
|
||||||
|
|
||||||
def read(self, size=-1):
|
def read(self, size: int = -1) -> bytes:
|
||||||
data = self._wrapped.read(size)
|
data = self._wrapped.read(size)
|
||||||
self._update(len(data))
|
self._update(len(data))
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
@contextlib.contextmanager
|
@contextlib.contextmanager
|
||||||
def file_progress(file, silent=False, **kwargs):
|
def file_progress(
|
||||||
|
file: io.IOBase, silent: bool = False, **kwargs: object
|
||||||
|
) -> Generator[Union[io.IOBase, "UpdateWrapper"], None, None]:
|
||||||
if silent:
|
if silent:
|
||||||
yield file
|
yield file
|
||||||
return
|
return
|
||||||
|
|
@ -206,8 +237,8 @@ def file_progress(file, silent=False, **kwargs):
|
||||||
if fileno == 0: # 0 means stdin
|
if fileno == 0: # 0 means stdin
|
||||||
yield file
|
yield file
|
||||||
else:
|
else:
|
||||||
file_length = os.path.getsize(file.name)
|
file_length = os.path.getsize(file.name) # type: ignore
|
||||||
with click.progressbar(length=file_length, **kwargs) as bar:
|
with click.progressbar(length=file_length, **kwargs) as bar: # type: ignore
|
||||||
yield UpdateWrapper(file, bar.update)
|
yield UpdateWrapper(file, bar.update)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -231,28 +262,30 @@ class RowError(Exception):
|
||||||
|
|
||||||
|
|
||||||
def _extra_key_strategy(
|
def _extra_key_strategy(
|
||||||
reader: Iterable[dict],
|
reader: Iterable[Dict[Optional[str], object]],
|
||||||
ignore_extras: Optional[bool] = False,
|
ignore_extras: Optional[bool] = False,
|
||||||
extras_key: Optional[str] = None,
|
extras_key: Optional[str] = None,
|
||||||
) -> Iterable[dict]:
|
) -> Iterable[Row]:
|
||||||
# Logic for handling CSV rows with more values than there are headings
|
# Logic for handling CSV rows with more values than there are headings
|
||||||
for row in reader:
|
for row in reader:
|
||||||
# DictReader adds a 'None' key with extra row values
|
# DictReader adds a 'None' key with extra row values
|
||||||
if None not in row:
|
if None not in row:
|
||||||
yield row
|
yield cast(Row, row)
|
||||||
elif ignore_extras:
|
elif ignore_extras:
|
||||||
# ignoring row.pop(none) because of this issue:
|
# ignoring row.pop(none) because of this issue:
|
||||||
# https://github.com/simonw/sqlite-utils/issues/440#issuecomment-1155358637
|
# https://github.com/simonw/sqlite-utils/issues/440#issuecomment-1155358637
|
||||||
row.pop(None) # type: ignore
|
row.pop(None)
|
||||||
yield row
|
yield cast(Row, row)
|
||||||
elif not extras_key:
|
elif not extras_key:
|
||||||
extras = row.pop(None) # type: ignore
|
extras = row.pop(None)
|
||||||
raise RowError(
|
raise RowError(
|
||||||
"Row {} contained these extra values: {}".format(row, extras)
|
"Row {} contained these extra values: {}".format(row, extras)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
row[extras_key] = row.pop(None) # type: ignore
|
extras_value = row.pop(None)
|
||||||
yield row
|
row_out = cast(Row, row)
|
||||||
|
row_out[extras_key] = extras_value # type: ignore[assignment]
|
||||||
|
yield row_out
|
||||||
|
|
||||||
|
|
||||||
def rows_from_file(
|
def rows_from_file(
|
||||||
|
|
@ -262,7 +295,7 @@ def rows_from_file(
|
||||||
encoding: Optional[str] = None,
|
encoding: Optional[str] = None,
|
||||||
ignore_extras: Optional[bool] = False,
|
ignore_extras: Optional[bool] = False,
|
||||||
extras_key: Optional[str] = None,
|
extras_key: Optional[str] = None,
|
||||||
) -> Tuple[Iterable[dict], Format]:
|
) -> Tuple[Iterable[Row], Format]:
|
||||||
"""
|
"""
|
||||||
Load a sequence of dictionaries from a file-like object containing one of four different formats.
|
Load a sequence of dictionaries from a file-like object containing one of four different formats.
|
||||||
|
|
||||||
|
|
@ -324,10 +357,17 @@ def rows_from_file(
|
||||||
rows = _extra_key_strategy(reader, ignore_extras, extras_key)
|
rows = _extra_key_strategy(reader, ignore_extras, extras_key)
|
||||||
return _CloseableIterator(iter(rows), decoded_fp), Format.CSV
|
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
|
||||||
)[0]
|
)
|
||||||
return _extra_key_strategy(rows, ignore_extras, extras_key), Format.TSV
|
return (
|
||||||
|
_extra_key_strategy(
|
||||||
|
cast(Iterable[Dict[Optional[str], object]], rows),
|
||||||
|
ignore_extras,
|
||||||
|
extras_key,
|
||||||
|
),
|
||||||
|
Format.TSV,
|
||||||
|
)
|
||||||
elif format is None:
|
elif format is None:
|
||||||
# Detect the format, then call this recursively
|
# Detect the format, then call this recursively
|
||||||
buffered = io.BufferedReader(cast(io.RawIOBase, fp), buffer_size=4096)
|
buffered = io.BufferedReader(cast(io.RawIOBase, fp), buffer_size=4096)
|
||||||
|
|
@ -349,8 +389,15 @@ def rows_from_file(
|
||||||
buffered, format=Format.CSV, dialect=dialect, encoding=encoding
|
buffered, format=Format.CSV, dialect=dialect, encoding=encoding
|
||||||
)
|
)
|
||||||
# Make sure we return the format we detected
|
# Make sure we return the format we detected
|
||||||
format = Format.TSV if dialect.delimiter == "\t" else Format.CSV
|
detected_format = Format.TSV if dialect.delimiter == "\t" else Format.CSV
|
||||||
return _extra_key_strategy(rows, ignore_extras, extras_key), format
|
return (
|
||||||
|
_extra_key_strategy(
|
||||||
|
cast(Iterable[Dict[Optional[str], object]], rows),
|
||||||
|
ignore_extras,
|
||||||
|
extras_key,
|
||||||
|
),
|
||||||
|
detected_format,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
raise RowsFromFileError("Bad format")
|
raise RowsFromFileError("Bad format")
|
||||||
|
|
||||||
|
|
@ -376,10 +423,10 @@ class TypeTracker:
|
||||||
db["creatures"].transform(types=tracker.types)
|
db["creatures"].transform(types=tracker.types)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self) -> None:
|
||||||
self.trackers = {}
|
self.trackers: Dict[str, "ValueTracker"] = {}
|
||||||
|
|
||||||
def wrap(self, iterator: Iterable[dict]) -> Iterable[dict]:
|
def wrap(self, iterator: Iterable[Dict[str, Any]]) -> Iterable[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Use this to loop through an existing iterator, tracking the column types
|
Use this to loop through an existing iterator, tracking the column types
|
||||||
as part of the iteration.
|
as part of the iteration.
|
||||||
|
|
@ -402,27 +449,29 @@ class TypeTracker:
|
||||||
|
|
||||||
|
|
||||||
class ValueTracker:
|
class ValueTracker:
|
||||||
def __init__(self):
|
couldbe: Dict[str, Callable[[object], bool]]
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
self.couldbe = {key: getattr(self, "test_" + key) for key in self.get_tests()}
|
self.couldbe = {key: getattr(self, "test_" + key) for key in self.get_tests()}
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_tests(cls):
|
def get_tests(cls) -> List[str]:
|
||||||
return [
|
return [
|
||||||
key.split("test_")[-1]
|
key.split("test_")[-1]
|
||||||
for key in cls.__dict__.keys()
|
for key in cls.__dict__.keys()
|
||||||
if key.startswith("test_")
|
if key.startswith("test_")
|
||||||
]
|
]
|
||||||
|
|
||||||
def test_integer(self, value):
|
def test_integer(self, value: object) -> bool:
|
||||||
try:
|
try:
|
||||||
int(value)
|
int(value) # type: ignore
|
||||||
return True
|
return True
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def test_float(self, value):
|
def test_float(self, value: object) -> bool:
|
||||||
try:
|
try:
|
||||||
float(value)
|
float(value) # type: ignore[arg-type]
|
||||||
return True
|
return True
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
return False
|
return False
|
||||||
|
|
@ -431,7 +480,7 @@ class ValueTracker:
|
||||||
return self.guessed_type + ": possibilities = " + repr(self.couldbe)
|
return self.guessed_type + ": possibilities = " + repr(self.couldbe)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def guessed_type(self):
|
def guessed_type(self) -> str:
|
||||||
options = set(self.couldbe.keys())
|
options = set(self.couldbe.keys())
|
||||||
# Return based on precedence
|
# Return based on precedence
|
||||||
for key in self.get_tests():
|
for key in self.get_tests():
|
||||||
|
|
@ -439,10 +488,10 @@ class ValueTracker:
|
||||||
return key
|
return key
|
||||||
return "text"
|
return "text"
|
||||||
|
|
||||||
def evaluate(self, value):
|
def evaluate(self, value: object) -> None:
|
||||||
if not value or not self.couldbe:
|
if not value or not self.couldbe:
|
||||||
return
|
return
|
||||||
not_these = []
|
not_these: List[str] = []
|
||||||
for name, test in self.couldbe.items():
|
for name, test in self.couldbe.items():
|
||||||
if not test(value):
|
if not test(value):
|
||||||
not_these.append(name)
|
not_these.append(name)
|
||||||
|
|
@ -451,45 +500,47 @@ class ValueTracker:
|
||||||
|
|
||||||
|
|
||||||
class NullProgressBar:
|
class NullProgressBar:
|
||||||
def __init__(self, *args):
|
def __init__(self, *args: Iterable[T]) -> None:
|
||||||
self.args = args
|
self.args = args
|
||||||
|
|
||||||
def __iter__(self):
|
def __iter__(self) -> Iterator[T]:
|
||||||
yield from self.args[0]
|
yield from self.args[0] # type: ignore
|
||||||
|
|
||||||
def update(self, value):
|
def update(self, value: int) -> None:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
@contextlib.contextmanager
|
@contextlib.contextmanager
|
||||||
def progressbar(*args, **kwargs):
|
def progressbar(*args: Iterable[T], **kwargs: Any) -> Generator[Any, None, None]:
|
||||||
silent = kwargs.pop("silent")
|
silent = kwargs.pop("silent")
|
||||||
if silent:
|
if silent:
|
||||||
yield NullProgressBar(*args)
|
yield NullProgressBar(*args)
|
||||||
else:
|
else:
|
||||||
with click.progressbar(*args, **kwargs) as bar:
|
with click.progressbar(*args, **kwargs) as bar: # type: ignore
|
||||||
yield bar
|
yield bar
|
||||||
|
|
||||||
|
|
||||||
def _compile_code(code, imports, variable="value"):
|
def _compile_code(
|
||||||
globals = {"r": recipes, "recipes": recipes}
|
code: str, imports: Iterable[str], variable: str = "value"
|
||||||
|
) -> Callable[..., Any]:
|
||||||
|
globals_dict: Dict[str, Any] = {"r": recipes, "recipes": recipes}
|
||||||
# Handle imports first so they're available for all approaches
|
# Handle imports first so they're available for all approaches
|
||||||
for import_ in imports:
|
for import_ in imports:
|
||||||
globals[import_.split(".")[0]] = __import__(import_)
|
globals_dict[import_.split(".")[0]] = __import__(import_)
|
||||||
|
|
||||||
# If user defined a convert() function, return that
|
# If user defined a convert() function, return that
|
||||||
try:
|
try:
|
||||||
exec(code, globals)
|
exec(code, globals_dict)
|
||||||
return globals["convert"]
|
return cast(Callable[..., object], globals_dict["convert"])
|
||||||
except (AttributeError, SyntaxError, NameError, KeyError, TypeError):
|
except (AttributeError, SyntaxError, NameError, KeyError, TypeError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# Check if code is a direct callable reference
|
# Check if code is a direct callable reference
|
||||||
# e.g. "r.parsedate" instead of "r.parsedate(value)"
|
# e.g. "r.parsedate" instead of "r.parsedate(value)"
|
||||||
try:
|
try:
|
||||||
fn = eval(code, globals)
|
fn = eval(code, globals_dict)
|
||||||
if callable(fn):
|
if callable(fn):
|
||||||
return fn
|
return cast(Callable[..., object], fn)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
@ -514,11 +565,11 @@ def _compile_code(code, imports, variable="value"):
|
||||||
if code_o is None:
|
if code_o is None:
|
||||||
raise SyntaxError("Could not compile code")
|
raise SyntaxError("Could not compile code")
|
||||||
|
|
||||||
exec(code_o, globals)
|
exec(code_o, globals_dict)
|
||||||
return globals["fn"]
|
return cast(Callable[..., object], globals_dict["fn"])
|
||||||
|
|
||||||
|
|
||||||
def chunks(sequence: Iterable, size: int) -> Iterable[Iterable]:
|
def chunks(sequence: Iterable[T], size: int) -> Iterable[Iterable[T]]:
|
||||||
"""
|
"""
|
||||||
Iterate over chunks of the sequence of the given size.
|
Iterate over chunks of the sequence of the given size.
|
||||||
|
|
||||||
|
|
@ -530,7 +581,7 @@ def chunks(sequence: Iterable, size: int) -> Iterable[Iterable]:
|
||||||
yield itertools.chain([item], itertools.islice(iterator, size - 1))
|
yield itertools.chain([item], itertools.islice(iterator, size - 1))
|
||||||
|
|
||||||
|
|
||||||
def hash_record(record: Dict, keys: Optional[Iterable[str]] = None):
|
def hash_record(record: Dict[str, Any], keys: Optional[Iterable[str]] = None) -> str:
|
||||||
"""
|
"""
|
||||||
``record`` should be a Python dictionary. Returns a sha1 hash of the
|
``record`` should be a Python dictionary. Returns a sha1 hash of the
|
||||||
keys and values in that record.
|
keys and values in that record.
|
||||||
|
|
@ -551,7 +602,7 @@ def hash_record(record: Dict, keys: Optional[Iterable[str]] = None):
|
||||||
:param record: Record to generate a hash for
|
:param record: Record to generate a hash for
|
||||||
:param keys: Subset of keys to use for that hash
|
:param keys: Subset of keys to use for that hash
|
||||||
"""
|
"""
|
||||||
to_hash = record
|
to_hash: Dict[str, Any] = record
|
||||||
if keys is not None:
|
if keys is not None:
|
||||||
to_hash = {key: record[key] for key in keys}
|
to_hash = {key: record[key] for key in keys}
|
||||||
return hashlib.sha1(
|
return hashlib.sha1(
|
||||||
|
|
@ -561,7 +612,7 @@ def hash_record(record: Dict, keys: Optional[Iterable[str]] = None):
|
||||||
).hexdigest()
|
).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
def _flatten(d):
|
def _flatten(d: Dict[str, Any]) -> Generator[Tuple[str, Any], None, None]:
|
||||||
for key, value in d.items():
|
for key, value in d.items():
|
||||||
if isinstance(value, dict):
|
if isinstance(value, dict):
|
||||||
for key2, value2 in _flatten(value):
|
for key2, value2 in _flatten(value):
|
||||||
|
|
@ -570,7 +621,7 @@ def _flatten(d):
|
||||||
yield key, value
|
yield key, value
|
||||||
|
|
||||||
|
|
||||||
def flatten(row: dict) -> dict:
|
def flatten(row: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Turn a nested dict e.g. ``{"a": {"b": 1}}`` into a flat dict: ``{"a_b": 1}``
|
Turn a nested dict e.g. ``{"a": {"b": 1}}`` into a flat dict: ``{"a_b": 1}``
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -41,9 +41,9 @@ def test_convert_help():
|
||||||
result = CliRunner().invoke(cli.cli, ["convert", "--help"])
|
result = CliRunner().invoke(cli.cli, ["convert", "--help"])
|
||||||
assert result.exit_code == 0
|
assert result.exit_code == 0
|
||||||
for expected in (
|
for expected in (
|
||||||
"r.jsonsplit(value, ",
|
"r.jsonsplit(value:",
|
||||||
"r.parsedate(value, ",
|
"r.parsedate(value:",
|
||||||
"r.parsedatetime(value, ",
|
"r.parsedatetime(value:",
|
||||||
):
|
):
|
||||||
assert expected in result.output
|
assert expected in result.output
|
||||||
|
|
||||||
|
|
@ -54,7 +54,7 @@ def test_convert_help():
|
||||||
n
|
n
|
||||||
for n in dir(recipes)
|
for n in dir(recipes)
|
||||||
if not n.startswith("_")
|
if not n.startswith("_")
|
||||||
and n not in ("json", "parser")
|
and n not in ("json", "parser", "Callable", "Optional")
|
||||||
and callable(getattr(recipes, n))
|
and callable(getattr(recipes, n))
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -50,7 +50,7 @@ def test_with_tracer():
|
||||||
with db.tracer(tracer):
|
with db.tracer(tracer):
|
||||||
list(dogs.search("Cleopaws"))
|
list(dogs.search("Cleopaws"))
|
||||||
|
|
||||||
assert len(collected) == 5
|
assert len(collected) == 4
|
||||||
assert collected == [
|
assert collected == [
|
||||||
(
|
(
|
||||||
"SELECT name FROM sqlite_master\n"
|
"SELECT name FROM sqlite_master\n"
|
||||||
|
|
@ -70,7 +70,6 @@ def test_with_tracer():
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
("select name from sqlite_master where type = 'view'", None),
|
("select name from sqlite_master where type = 'view'", None),
|
||||||
("select name from sqlite_master where type = 'view'", None),
|
|
||||||
("select sql from sqlite_master where name = ?", ("dogs_fts",)),
|
("select sql from sqlite_master where name = ?", ("dogs_fts",)),
|
||||||
(
|
(
|
||||||
'with "original" as (\n'
|
'with "original" as (\n'
|
||||||
|
|
@ -94,4 +93,4 @@ def test_with_tracer():
|
||||||
|
|
||||||
# Outside the with block collected should not be appended to
|
# Outside the with block collected should not be appended to
|
||||||
dogs.insert({"name": "Cleopaws"})
|
dogs.insert({"name": "Cleopaws"})
|
||||||
assert len(collected) == 5
|
assert len(collected) == 4
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue