diff --git a/pyproject.toml b/pyproject.toml index f6dd0a8..5bfa492 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -47,7 +47,7 @@ dev = [ # flake8 "flake8", "flake8-pyproject", - "ty", + "ty>=0.0.37", # For stable cog: "tabulate>=0.10.0", ] diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index cbec3b2..5844dfc 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -1,7 +1,7 @@ import base64 from typing import Any import click -from click_default_group import DefaultGroup # type: ignore +from click_default_group import DefaultGroup from datetime import datetime, timezone import hashlib import pathlib @@ -222,7 +222,7 @@ def tables( else: items = db.table_names(fts4=fts4, fts5=fts5) for name in items: - row = [name] + row: list[Any] = [name] if counts: row.append(method(name).count) if columns: @@ -2121,8 +2121,9 @@ def _execute_query( cursor = [[cursor.rowcount]] else: headers = [c[0] for c in cursor.description] + cursor_or_rows: Any = cursor if raw: - row = cursor.fetchone() # type: ignore[union-attr] + row = cursor_or_rows.fetchone() data = row[0] if row else None if isinstance(data, bytes): sys.stdout.buffer.write(data) diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 5a3dcf1..ed3fc7a 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -15,6 +15,7 @@ from collections.abc import Mapping import contextlib import datetime import decimal +import importlib import inspect import itertools import json @@ -22,7 +23,7 @@ import os import pathlib import re import secrets -from sqlite_fts4 import rank_bm25 # type: ignore +from sqlite_fts4 import rank_bm25 import textwrap from typing import ( cast, @@ -43,7 +44,7 @@ import uuid from sqlite_utils.plugins import pm try: - from sqlite_dump import iterdump # type: ignore[import-not-found] + iterdump = importlib.import_module("sqlite_dump").iterdump except ImportError: iterdump = None @@ -82,15 +83,17 @@ def quote_identifier(identifier: str) -> str: return '"{}"'.format(identifier.replace('"', '""')) +pd: Any = None try: - import pandas as pd # type: ignore + pd = importlib.import_module("pandas") except ImportError: - pd = None # type: ignore + pd = None +np: Any = None try: - import numpy as np # type: ignore + np = importlib.import_module("numpy") except ImportError: - np = None # type: ignore + np = None Column = namedtuple( "Column", ("cid", "name", "type", "notnull", "default_value", "is_pk") @@ -190,7 +193,10 @@ class Default: DEFAULT = Default() -COLUMN_TYPE_MAPPING = { +Tracer = Callable[[str, Optional[Union[Sequence[Any], Dict[str, Any]]]], None] + + +COLUMN_TYPE_MAPPING: Dict[Any, str] = { float: "REAL", int: "INTEGER", bool: "INTEGER", @@ -339,7 +345,7 @@ class Database: memory_name: Optional[str] = None, recreate: bool = False, recursive_triggers: bool = True, - tracer: Optional[Callable] = None, + tracer: Optional[Tracer] = None, use_counts_table: bool = False, execute_plugins: bool = True, use_old_upsert: bool = False, @@ -375,8 +381,8 @@ class Database: self.conn = sqlite3.connect(str(filename_or_conn)) else: assert not recreate, "recreate cannot be used with connections, only paths" - self.conn = filename_or_conn - self._tracer = tracer + self.conn = cast(sqlite3.Connection, filename_or_conn) + self._tracer: Optional[Tracer] = tracer if recursive_triggers: self.execute("PRAGMA recursive_triggers=on;") self._registered_functions: set = set() @@ -421,7 +427,7 @@ class Database: @contextlib.contextmanager def tracer( - self, tracer: Optional[Callable[[str, Optional[Sequence]], None]] = None + self, tracer: Optional[Tracer] = None ) -> Generator["Database", None, None]: """ Context manager to temporarily set a tracer function - all executed SQL queries will @@ -439,7 +445,7 @@ class Database: :param tracer: Callable accepting ``sql`` and ``parameters`` arguments """ prev_tracer = self._tracer - self._tracer = tracer or print + self._tracer = tracer or cast(Tracer, print) try: yield self finally: @@ -3493,7 +3499,7 @@ class Table(Queryable): raise ValueError( "When using list-based iteration, the first yielded value must be a list of column name strings" ) - column_names = list(first_record) + column_names = cast(List[str], list(first_record)) all_columns = column_names num_columns = len(column_names) # Get the actual first data record @@ -3535,7 +3541,8 @@ class Table(Queryable): chunk_as_dicts = [dict(zip(column_names, row)) for row in chunk] column_types = suggest_column_types(chunk_as_dicts) else: - column_types = suggest_column_types(chunk) # type: ignore[arg-type] + dict_chunk = cast(List[Dict[str, Any]], chunk) + column_types = suggest_column_types(dict_chunk) if extracts: for col in extracts: if col in column_types: @@ -3562,14 +3569,14 @@ class Table(Queryable): all_columns.insert(0, hash_id) else: all_columns_set: Set[str] = set() - for record in chunk: - all_columns_set.update(record.keys()) # type: ignore[union-attr] + for record in cast(List[Dict[str, Any]], chunk): + all_columns_set.update(record.keys()) all_columns = list(sorted(all_columns_set)) if hash_id: all_columns.insert(0, hash_id) else: if not list_mode: - for record in chunk: + for record in cast(List[Dict[str, Any]], chunk): all_columns += [ column for column in record if column not in all_columns ] @@ -3767,6 +3774,7 @@ class Table(Queryable): :param strict: Boolean, apply STRICT mode if creating the table. """ assert isinstance(lookup_values, dict) + assert pk is not None if extra_values is not None: assert isinstance(extra_values, dict) combined_values = dict(lookup_values) @@ -3786,7 +3794,7 @@ class Table(Queryable): ) ) try: - return rows[0][pk] # type: ignore[index] + return rows[0][pk] except IndexError: return self.insert( combined_values, diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index 05e9a51..0ca98fc 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -3,6 +3,7 @@ import contextlib import csv import enum import hashlib +import importlib import io import itertools import json @@ -21,6 +22,7 @@ from typing import ( Set, Tuple, Type, + TYPE_CHECKING, TypeVar, Union, cast, @@ -30,22 +32,26 @@ import click from . import recipes -try: - import pysqlite3 as sqlite3 # type: ignore[import-not-found] # noqa: F401 - from pysqlite3 import dbapi2 # type: ignore[import-not-found] # noqa: F401 +if TYPE_CHECKING: + import sqlite3 # noqa: F401 + from sqlite3 import dbapi2 # noqa: F401 OperationalError = dbapi2.OperationalError -except ImportError: +else: try: - import sqlean as sqlite3 # type: ignore[import-not-found] # noqa: F401 - from sqlean import dbapi2 # type: ignore[import-not-found] # noqa: F401 - + sqlite3 = importlib.import_module("pysqlite3") + dbapi2 = importlib.import_module("pysqlite3.dbapi2") OperationalError = dbapi2.OperationalError except ImportError: - import sqlite3 # noqa: F401 - from sqlite3 import dbapi2 # noqa: F401 + try: + sqlite3 = importlib.import_module("sqlean") + dbapi2 = importlib.import_module("sqlean.dbapi2") + OperationalError = dbapi2.OperationalError + except ImportError: + import sqlite3 # noqa: F401 + from sqlite3 import dbapi2 # noqa: F401 - OperationalError = dbapi2.OperationalError + OperationalError = dbapi2.OperationalError SPATIALITE_PATHS = ( @@ -60,7 +66,7 @@ SPATIALITE_PATHS = ( 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] +RowValue = Union[None, int, float, str, bytes, bool, List[str]] Row = Dict[str, RowValue] T = TypeVar("T") @@ -284,7 +290,7 @@ def _extra_key_strategy( else: extras_value = row.pop(None) row_out = cast(Row, row) - row_out[extras_key] = extras_value # type: ignore[assignment] + row_out[extras_key] = cast(RowValue, extras_value) yield row_out @@ -464,14 +470,14 @@ class ValueTracker: def test_integer(self, value: object) -> bool: try: - int(value) # type: ignore + int(cast(Any, value)) return True except (ValueError, TypeError): return False def test_float(self, value: object) -> bool: try: - float(value) # type: ignore[arg-type] + float(cast(Any, value)) return True except (ValueError, TypeError): return False