mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-07-31 22:44:11 +02:00
Add type hints to public APIs and configure mypy
- Improve mypy configuration with appropriate strictness settings
- Add comprehensive type hints to public API functions:
- hookspecs.py: Type register_commands and prepare_connection hooks
- plugins.py: Type get_plugins function and PluginManager
- recipes.py: Type parsedate, parsedatetime, jsonsplit functions
- utils.py: Type all public utility functions including rows_from_file,
TypeTracker, ValueTracker, chunks, hash_record, flatten, etc.
- db.py: Type Database, Queryable, Table, and View class methods
including constructor, execute, query, table operations, etc.
- Exclude cli.py from strict checking (internal implementation detail)
- All public API modules now pass mypy type checking
This commit is contained in:
parent
fd5b09f64b
commit
83c50527aa
6 changed files with 226 additions and 126 deletions
|
|
@ -8,7 +8,22 @@ import itertools
|
|||
import json
|
||||
import os
|
||||
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,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
import click
|
||||
|
||||
|
|
@ -44,17 +59,19 @@ SPATIALITE_PATHS = (
|
|||
ORIGINAL_CSV_FIELD_SIZE_LIMIT = csv.field_size_limit()
|
||||
|
||||
|
||||
class _CloseableIterator(Iterator[dict]):
|
||||
class _CloseableIterator(Iterator[Dict[str, Any]]):
|
||||
"""Iterator wrapper that closes a file when iteration is complete."""
|
||||
|
||||
def __init__(self, iterator: Iterator[dict], closeable: io.IOBase):
|
||||
def __init__(
|
||||
self, iterator: Iterator[Dict[str, Any]], closeable: io.IOBase
|
||||
) -> None:
|
||||
self._iterator = iterator
|
||||
self._closeable = closeable
|
||||
|
||||
def __iter__(self) -> "_CloseableIterator":
|
||||
return self
|
||||
|
||||
def __next__(self) -> dict:
|
||||
def __next__(self) -> Dict[str, Any]:
|
||||
try:
|
||||
return next(self._iterator)
|
||||
except StopIteration:
|
||||
|
|
@ -65,7 +82,7 @@ class _CloseableIterator(Iterator[dict]):
|
|||
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.
|
||||
"""
|
||||
|
|
@ -108,20 +125,25 @@ def find_spatialite() -> Optional[str]:
|
|||
return None
|
||||
|
||||
|
||||
def suggest_column_types(records):
|
||||
all_column_types = {}
|
||||
def suggest_column_types(
|
||||
records: Iterable[Dict[str, Any]],
|
||||
) -> Dict[str, type]:
|
||||
all_column_types: Dict[str, Set[type]] = {}
|
||||
for record in records:
|
||||
for key, value in record.items():
|
||||
all_column_types.setdefault(key, set()).add(type(value))
|
||||
return types_for_column_types(all_column_types)
|
||||
|
||||
|
||||
def types_for_column_types(all_column_types):
|
||||
column_types = {}
|
||||
def types_for_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():
|
||||
# Ignore null values if at least one other type present:
|
||||
if len(types) > 1:
|
||||
types.discard(None.__class__)
|
||||
t: type
|
||||
if {None.__class__} == types:
|
||||
t = str
|
||||
elif len(types) == 1:
|
||||
|
|
@ -143,7 +165,7 @@ def types_for_column_types(all_column_types):
|
|||
return column_types
|
||||
|
||||
|
||||
def column_affinity(column_type):
|
||||
def column_affinity(column_type: str) -> type:
|
||||
# Implementation of SQLite affinity rules from
|
||||
# https://www.sqlite.org/datatype3.html#determination_of_column_affinity
|
||||
assert isinstance(column_type, str)
|
||||
|
|
@ -162,7 +184,7 @@ def column_affinity(column_type):
|
|||
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
|
||||
to_fix = [
|
||||
k
|
||||
|
|
@ -177,23 +199,25 @@ def decode_base64_values(doc):
|
|||
|
||||
|
||||
class UpdateWrapper:
|
||||
def __init__(self, wrapped, update):
|
||||
def __init__(self, wrapped: Any, update: Callable[[int], None]) -> None:
|
||||
self._wrapped = wrapped
|
||||
self._update = update
|
||||
|
||||
def __iter__(self):
|
||||
def __iter__(self) -> Iterator[Any]:
|
||||
for line in self._wrapped:
|
||||
self._update(len(line))
|
||||
yield line
|
||||
|
||||
def read(self, size=-1):
|
||||
def read(self, size: int = -1) -> Any:
|
||||
data = self._wrapped.read(size)
|
||||
self._update(len(data))
|
||||
return data
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def file_progress(file, silent=False, **kwargs):
|
||||
def file_progress(
|
||||
file: Any, silent: bool = False, **kwargs: Any
|
||||
) -> Generator[Any, None, None]:
|
||||
if silent:
|
||||
yield file
|
||||
return
|
||||
|
|
@ -231,28 +255,30 @@ class RowError(Exception):
|
|||
|
||||
|
||||
def _extra_key_strategy(
|
||||
reader: Iterable[dict],
|
||||
reader: Iterable[Dict[Any, Any]],
|
||||
ignore_extras: Optional[bool] = False,
|
||||
extras_key: Optional[str] = None,
|
||||
) -> Iterable[dict]:
|
||||
) -> Iterable[Dict[str, Any]]:
|
||||
# Logic for handling CSV rows with more values than there are headings
|
||||
for row in reader:
|
||||
# DictReader adds a 'None' key with extra row values
|
||||
if None not in row:
|
||||
yield row
|
||||
row_with_optional_none = cast(Dict[Optional[str], Any], row)
|
||||
if None not in row_with_optional_none:
|
||||
yield cast(Dict[str, Any], row)
|
||||
elif ignore_extras:
|
||||
# ignoring row.pop(none) because of this issue:
|
||||
# https://github.com/simonw/sqlite-utils/issues/440#issuecomment-1155358637
|
||||
row.pop(None) # type: ignore
|
||||
yield row
|
||||
row_with_optional_none.pop(None)
|
||||
yield cast(Dict[str, Any], row)
|
||||
elif not extras_key:
|
||||
extras = row.pop(None) # type: ignore
|
||||
extras = row_with_optional_none.pop(None)
|
||||
raise RowError(
|
||||
"Row {} contained these extra values: {}".format(row, extras)
|
||||
)
|
||||
else:
|
||||
row[extras_key] = row.pop(None) # type: ignore
|
||||
yield row
|
||||
row_out = cast(Dict[str, Any], row)
|
||||
row_out[extras_key] = row_with_optional_none.pop(None)
|
||||
yield row_out
|
||||
|
||||
|
||||
def rows_from_file(
|
||||
|
|
@ -262,7 +288,7 @@ def rows_from_file(
|
|||
encoding: Optional[str] = None,
|
||||
ignore_extras: Optional[bool] = False,
|
||||
extras_key: Optional[str] = None,
|
||||
) -> Tuple[Iterable[dict], Format]:
|
||||
) -> Tuple[Iterable[Dict[str, Any]], Format]:
|
||||
"""
|
||||
Load a sequence of dictionaries from a file-like object containing one of four different formats.
|
||||
|
||||
|
|
@ -376,10 +402,10 @@ class TypeTracker:
|
|||
db["creatures"].transform(types=tracker.types)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.trackers = {}
|
||||
def __init__(self) -> None:
|
||||
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
|
||||
as part of the iteration.
|
||||
|
|
@ -402,25 +428,27 @@ class TypeTracker:
|
|||
|
||||
|
||||
class ValueTracker:
|
||||
def __init__(self):
|
||||
couldbe: Dict[str, Callable[[Any], bool]]
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.couldbe = {key: getattr(self, "test_" + key) for key in self.get_tests()}
|
||||
|
||||
@classmethod
|
||||
def get_tests(cls):
|
||||
def get_tests(cls) -> List[str]:
|
||||
return [
|
||||
key.split("test_")[-1]
|
||||
for key in cls.__dict__.keys()
|
||||
if key.startswith("test_")
|
||||
]
|
||||
|
||||
def test_integer(self, value):
|
||||
def test_integer(self, value: Any) -> bool:
|
||||
try:
|
||||
int(value)
|
||||
return True
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
|
||||
def test_float(self, value):
|
||||
def test_float(self, value: Any) -> bool:
|
||||
try:
|
||||
float(value)
|
||||
return True
|
||||
|
|
@ -431,7 +459,7 @@ class ValueTracker:
|
|||
return self.guessed_type + ": possibilities = " + repr(self.couldbe)
|
||||
|
||||
@property
|
||||
def guessed_type(self):
|
||||
def guessed_type(self) -> str:
|
||||
options = set(self.couldbe.keys())
|
||||
# Return based on precedence
|
||||
for key in self.get_tests():
|
||||
|
|
@ -439,10 +467,10 @@ class ValueTracker:
|
|||
return key
|
||||
return "text"
|
||||
|
||||
def evaluate(self, value):
|
||||
def evaluate(self, value: Any) -> None:
|
||||
if not value or not self.couldbe:
|
||||
return
|
||||
not_these = []
|
||||
not_these: List[str] = []
|
||||
for name, test in self.couldbe.items():
|
||||
if not test(value):
|
||||
not_these.append(name)
|
||||
|
|
@ -451,18 +479,18 @@ class ValueTracker:
|
|||
|
||||
|
||||
class NullProgressBar:
|
||||
def __init__(self, *args):
|
||||
def __init__(self, *args: Any) -> None:
|
||||
self.args = args
|
||||
|
||||
def __iter__(self):
|
||||
def __iter__(self) -> Iterator[Any]:
|
||||
yield from self.args[0]
|
||||
|
||||
def update(self, value):
|
||||
def update(self, value: int) -> None:
|
||||
pass
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def progressbar(*args, **kwargs):
|
||||
def progressbar(*args: Any, **kwargs: Any) -> Generator[Any, None, None]:
|
||||
silent = kwargs.pop("silent")
|
||||
if silent:
|
||||
yield NullProgressBar(*args)
|
||||
|
|
@ -471,25 +499,27 @@ def progressbar(*args, **kwargs):
|
|||
yield bar
|
||||
|
||||
|
||||
def _compile_code(code, imports, variable="value"):
|
||||
globals = {"r": recipes, "recipes": recipes}
|
||||
def _compile_code(
|
||||
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
|
||||
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
|
||||
try:
|
||||
exec(code, globals)
|
||||
return globals["convert"]
|
||||
exec(code, globals_dict)
|
||||
return cast(Callable[..., Any], globals_dict["convert"])
|
||||
except (AttributeError, SyntaxError, NameError, KeyError, TypeError):
|
||||
pass
|
||||
|
||||
# Check if code is a direct callable reference
|
||||
# e.g. "r.parsedate" instead of "r.parsedate(value)"
|
||||
try:
|
||||
fn = eval(code, globals)
|
||||
fn = eval(code, globals_dict)
|
||||
if callable(fn):
|
||||
return fn
|
||||
return cast(Callable[..., Any], fn)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -514,11 +544,11 @@ def _compile_code(code, imports, variable="value"):
|
|||
if code_o is None:
|
||||
raise SyntaxError("Could not compile code")
|
||||
|
||||
exec(code_o, globals)
|
||||
return globals["fn"]
|
||||
exec(code_o, globals_dict)
|
||||
return cast(Callable[..., Any], globals_dict["fn"])
|
||||
|
||||
|
||||
def chunks(sequence: Iterable, size: int) -> Iterable[Iterable]:
|
||||
def chunks(sequence: Iterable[Any], size: int) -> Iterable[Iterable[Any]]:
|
||||
"""
|
||||
Iterate over chunks of the sequence of the given size.
|
||||
|
||||
|
|
@ -530,7 +560,9 @@ def chunks(sequence: Iterable, size: int) -> Iterable[Iterable]:
|
|||
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
|
||||
keys and values in that record.
|
||||
|
|
@ -551,7 +583,7 @@ def hash_record(record: Dict, keys: Optional[Iterable[str]] = None):
|
|||
:param record: Record to generate a hash for
|
||||
:param keys: Subset of keys to use for that hash
|
||||
"""
|
||||
to_hash = record
|
||||
to_hash: Dict[str, Any] = record
|
||||
if keys is not None:
|
||||
to_hash = {key: record[key] for key in keys}
|
||||
return hashlib.sha1(
|
||||
|
|
@ -561,7 +593,7 @@ def hash_record(record: Dict, keys: Optional[Iterable[str]] = None):
|
|||
).hexdigest()
|
||||
|
||||
|
||||
def _flatten(d):
|
||||
def _flatten(d: Dict[str, Any]) -> Generator[Tuple[str, Any], None, None]:
|
||||
for key, value in d.items():
|
||||
if isinstance(value, dict):
|
||||
for key2, value2 in _flatten(value):
|
||||
|
|
@ -570,7 +602,7 @@ def _flatten(d):
|
|||
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}``
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue