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:
Simon Willison 2025-12-16 22:11:47 -08:00 committed by GitHub
commit 8d74ffc932
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 274 additions and 154 deletions

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -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]

View file

@ -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"""

View file

@ -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],
} }

View file

@ -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"]
""" """

View file

@ -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}``

View file

@ -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))
], ],
) )

View file

@ -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