diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index 54de265..c6e70fe 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -1051,7 +1051,7 @@ def insert_upsert_implementation( csv_reader_args["delimiter"] = delimiter if 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) if no_headers: headers = ["untitled_{}".format(i + 1) for i in range(len(first_row))] diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index d855ad3..e3af8bc 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -489,7 +489,7 @@ class Database: """ 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) if not replace and (fn_name, arity) in self._registered_functions: return fn @@ -1450,11 +1450,11 @@ class Queryable: def pks_and_rows_where( self, 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, limit: 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. @@ -1848,7 +1848,7 @@ class Table(Queryable): quote_identifier(self.name), ) self.db.execute(sql) - return self.db[new_name] + return self.db.table(new_name) def transform( self, @@ -2153,7 +2153,7 @@ class Table(Queryable): ) ) table = table or "_".join(columns) - lookup_table = self.db[table] + lookup_table = self.db.table(table) fk_column = fk_column or "{}_id".format(table) magic_lookup_column = "{}_{}".format(fk_column, os.urandom(6).hex()) @@ -2680,7 +2680,7 @@ class Table(Queryable): ) return self - def rebuild_fts(self) -> None: + def rebuild_fts(self) -> "Table": "Run the ``rebuild`` operation against the associated full-text search index table." fts_table = self.detect_fts() if fts_table is None: @@ -2767,7 +2767,7 @@ class Table(Queryable): self.name ) 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( """ with {original} as ( @@ -2887,7 +2887,7 @@ class Table(Queryable): def delete_where( self, where: Optional[str] = None, - where_args: Optional[Union[Iterable, dict]] = None, + where_args: Optional[Union[Sequence, Dict[str, Any]]] = None, analyze: bool = False, ) -> "Table": """ @@ -2978,9 +2978,9 @@ class Table(Queryable): drop: bool = False, multi: bool = False, where: Optional[str] = None, - where_args: Optional[Union[Iterable, dict]] = None, + where_args: Optional[Union[Sequence, Dict[str, Any]]] = None, show_progress: bool = False, - ): + ) -> "Table": """ Apply conversion function ``fn`` to every value in the specified columns. @@ -3143,7 +3143,7 @@ class Table(Queryable): if has_extracts: for i, key in enumerate(all_columns): 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]} ) values.append(record_values) @@ -3164,7 +3164,7 @@ class Table(Queryable): ) if key in extracts: 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) values.append(record_values) @@ -3874,7 +3874,7 @@ class Table(Queryable): already exists. """ 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 if lookup is not None: assert record_or_iterable is None, "Provide lookup= or record, not both" diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index cdc1ac1..06b4299 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -237,8 +237,8 @@ def file_progress( if fileno == 0: # 0 means stdin yield file else: - file_length = os.path.getsize(file.name) - with click.progressbar(length=file_length, **kwargs) as bar: + file_length = os.path.getsize(file.name) # type: ignore + with click.progressbar(length=file_length, **kwargs) as bar: # type: ignore yield UpdateWrapper(file, bar.update) @@ -516,7 +516,7 @@ def progressbar(*args: Iterable[T], **kwargs: Any) -> Generator[Any, None, None] if silent: yield NullProgressBar(*args) else: - with click.progressbar(*args, **kwargs) as bar: + with click.progressbar(*args, **kwargs) as bar: # type: ignore yield bar diff --git a/tests/test_tracer.py b/tests/test_tracer.py index 9b44fce..ac490c5 100644 --- a/tests/test_tracer.py +++ b/tests/test_tracer.py @@ -50,7 +50,7 @@ def test_with_tracer(): with db.tracer(tracer): list(dogs.search("Cleopaws")) - assert len(collected) == 5 + assert len(collected) == 4 assert collected == [ ( "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 sql from sqlite_master where name = ?", ("dogs_fts",)), ( 'with "original" as (\n' @@ -94,4 +93,4 @@ def test_with_tracer(): # Outside the with block collected should not be appended to dogs.insert({"name": "Cleopaws"}) - assert len(collected) == 5 + assert len(collected) == 4