From d71420065903ff54247b5062b8c6af6165b7e638 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sun, 12 Jul 2026 05:00:32 -0700 Subject: [PATCH 1/6] Add test: transform does not cascade-delete referencing records (#792) > Add a test that covers what happens if you run transform against a table with ON CASCADE DELETE for one of its foreign keys - those records should not be deleted during the transform even though the table is dropped as part of that procedure --- tests/test_transform.py | 37 +++++++++++++++++++++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/tests/test_transform.py b/tests/test_transform.py index f0f5019..3cc5ba6 100644 --- a/tests/test_transform.py +++ b/tests/test_transform.py @@ -432,6 +432,43 @@ def test_transform_verify_foreign_keys(fresh_db): assert fresh_db.conn.execute("PRAGMA foreign_keys").fetchone()[0] +@pytest.mark.parametrize("use_pragma_foreign_keys", [False, True]) +def test_transform_on_delete_cascade_does_not_delete_records( + fresh_db, use_pragma_foreign_keys +): + # Transforming a table drops and recreates it - if another table references + # it with ON DELETE CASCADE and PRAGMA foreign_keys is on, that drop must + # not cascade and delete the referencing records + if use_pragma_foreign_keys: + fresh_db.conn.execute("PRAGMA foreign_keys=ON") + fresh_db.executescript(""" + CREATE TABLE authors (id INTEGER PRIMARY KEY, name TEXT); + CREATE TABLE books ( + id INTEGER PRIMARY KEY, + title TEXT, + author_id INTEGER REFERENCES authors(id) ON DELETE CASCADE + ); + """) + fresh_db["authors"].insert({"id": 1, "name": "Ursula K. Le Guin"}) + fresh_db["books"].insert({"id": 1, "title": "The Dispossessed", "author_id": 1}) + # Transform the table on the other end of the cascading foreign key + fresh_db["authors"].transform(rename={"name": "author_name"}) + assert list(fresh_db["authors"].rows) == [ + {"id": 1, "author_name": "Ursula K. Le Guin"} + ] + assert list(fresh_db["books"].rows) == [ + {"id": 1, "title": "The Dispossessed", "author_id": 1} + ] + # Transforming the table with the cascading foreign key should not + # delete its records either + fresh_db["books"].transform(rename={"title": "book_title"}) + assert list(fresh_db["books"].rows) == [ + {"id": 1, "book_title": "The Dispossessed", "author_id": 1} + ] + if use_pragma_foreign_keys: + assert fresh_db.conn.execute("PRAGMA foreign_keys").fetchone()[0] + + def test_transform_add_foreign_keys_from_scratch(fresh_db): _add_country_city_continent(fresh_db) fresh_db["places"].insert(_CAVEAU) From f66ddcb215e76dcbc1fa1dff4359f2ac3dc702e5 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sun, 12 Jul 2026 08:43:51 -0700 Subject: [PATCH 2/6] Transform now refuses to run inside a transaction if destructive foreign keys exist (#795) * Transform now refuses to run inside a transaction if destructive foreign keys exist Closes #794 Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_014StVTWQJpFhfZJK2CYVBwv --- docs/python-api.rst | 33 ++++++++++- sqlite_utils/db.py | 35 ++++++++++++ tests/test_transform.py | 124 +++++++++++++++++++++++++++++++++++++++- 3 files changed, 190 insertions(+), 2 deletions(-) diff --git a/docs/python-api.rst b/docs/python-api.rst index 1ed238e..43b734d 100644 --- a/docs/python-api.rst +++ b/docs/python-api.rst @@ -434,9 +434,10 @@ The library will never commit a transaction you opened. If you call write method Prefer ``db.atomic()`` or ``db.begin()``, ``db.commit()`` and ``db.rollback()`` over mixing sqlite-utils transaction methods with calls to ``db.conn.commit()``, ``db.conn.rollback()`` or raw transaction-control SQL. Mixing the two layers makes it much harder to tell which layer owns the current transaction. -Two related safeguards to be aware of: +Some related safeguards to be aware of: - ``db.enable_wal()`` and ``db.disable_wal()`` raise a ``sqlite_utils.db.TransactionError`` if called while a transaction is open, because changing the journal mode would commit it as a side effect. +- ``table.transform()`` raises a ``sqlite_utils.db.TransactionError`` if called while a transaction is open with ``PRAGMA foreign_keys`` enabled and the table is referenced by foreign keys with destructive ``ON DELETE`` actions, because the pragma cannot be turned off mid-transaction to protect those referencing rows - see :ref:`python_api_transform_foreign_keys_transactions`. - Closing the database - explicitly with ``db.close()``, or by exiting a ``with Database(...) as db:`` block - rolls back any transaction that is still open, see :ref:`python_api_close`. .. _python_api_transactions_modes: @@ -1996,6 +1997,36 @@ If you want to do something more advanced, you can call the ``table.transform_sq This method will return a list of SQL statements that should be executed to implement the change. You can then make modifications to that SQL - or add additional SQL statements - before executing it yourself. +.. _python_api_transform_foreign_keys_transactions: + +Foreign keys and transactions +----------------------------- + +Because ``.transform()`` drops the old table, running it with ``PRAGMA foreign_keys`` enabled could fire ``ON DELETE`` actions on any tables that reference it - an inbound ``ON DELETE CASCADE`` foreign key would silently delete those referencing rows. To prevent this, ``.transform()`` turns ``PRAGMA foreign_keys`` off for the duration of the operation and restores it afterwards, running ``PRAGMA foreign_key_check`` before committing. + +``PRAGMA foreign_keys`` cannot be changed inside a transaction, so this protection is impossible if you call ``.transform()`` while a transaction is already open - for example inside a ``with db.atomic():`` block or after ``db.begin()``. If ``PRAGMA foreign_keys`` is on and another table references the table being transformed with a destructive ``ON DELETE`` action - ``CASCADE``, ``SET NULL`` or ``SET DEFAULT`` - the method will refuse to run and raise a ``sqlite_utils.db.TransactionError``: + +.. code-block:: python + + from sqlite_utils.db import TransactionError + + try: + with db.atomic(): + db["authors"].transform(types={"id": str}) + except TransactionError as ex: + print("Could not transform in transaction:", ex) + +To transform such a table either call ``.transform()`` outside of the transaction, or execute ``PRAGMA foreign_keys = off`` before opening it: + +.. code-block:: python + + db.execute("PRAGMA foreign_keys = off") + with db.atomic(): + db["authors"].transform(types={"id": str}) + db.execute("PRAGMA foreign_keys = on") + +Tables referenced by foreign keys without a destructive action (the default ``NO ACTION``, or ``RESTRICT``) can still be transformed inside a transaction - sqlite-utils uses ``PRAGMA defer_foreign_keys`` to postpone the foreign key checks until the transaction commits. + .. _python_api_extract: Extracting columns into a separate table diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index d709fb9..e97b7d9 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -2522,6 +2522,11 @@ class Table(Queryable): See :ref:`python_api_transform` for full details. + Raises :py:class:`sqlite_utils.db.TransactionError` if called while a + transaction is open with ``PRAGMA foreign_keys`` enabled and the table + is referenced by foreign keys with destructive ``ON DELETE`` actions - + see :ref:`python_api_transform_foreign_keys_transactions`. + :param types: Columns that should have their type changed, for example ``{"weight": float}`` :param rename: Columns to rename, for example ``{"headline": "title"}`` :param drop: Columns to drop @@ -2566,6 +2571,36 @@ class Table(Queryable): should_defer_foreign_keys = ( pragma_foreign_keys_was_on and already_in_transaction ) + if should_defer_foreign_keys: + # PRAGMA foreign_keys is a no-op inside a transaction, and + # defer_foreign_keys only defers violation checks, not ON DELETE + # actions - so dropping the old table would still fire destructive + # actions on any tables that reference it. Refuse rather than + # silently modify or delete those rows. + destructive_fks = [ + (table.name, fk) + for table in self.db.tables + for fk in table.foreign_keys + if fk.other_table == self.name + and fk.on_delete in ("CASCADE", "SET NULL", "SET DEFAULT") + ] + if destructive_fks: + raise TransactionError( + "Cannot transform table {table} while a transaction is open: " + "PRAGMA foreign_keys cannot be changed inside a transaction, " + "and the table is referenced by foreign keys with ON DELETE " + "actions that would fire when the old table is dropped: " + "{fks}. Call transform() outside of the transaction, or " + 'execute "PRAGMA foreign_keys = off" before opening it.'.format( + table=self.name, + fks=", ".join( + "{}.{} (ON DELETE {})".format( + table_name, ", ".join(fk.columns), fk.on_delete + ) + for table_name, fk in destructive_fks + ), + ) + ) defer_foreign_keys_was_on = False try: if should_disable_foreign_keys: diff --git a/tests/test_transform.py b/tests/test_transform.py index 3cc5ba6..362f1ca 100644 --- a/tests/test_transform.py +++ b/tests/test_transform.py @@ -1,6 +1,6 @@ import sqlite3 -from sqlite_utils.db import ForeignKey, TransformError +from sqlite_utils.db import ForeignKey, TransactionError, TransformError from sqlite_utils.utils import OperationalError import pytest @@ -469,6 +469,128 @@ def test_transform_on_delete_cascade_does_not_delete_records( assert fresh_db.conn.execute("PRAGMA foreign_keys").fetchone()[0] +@pytest.mark.parametrize("on_delete", ["CASCADE", "SET NULL", "SET DEFAULT", "cascade"]) +def test_transform_in_transaction_refuses_destructive_on_delete(fresh_db, on_delete): + # PRAGMA foreign_keys is a no-op inside a transaction, so transforming a + # table referenced by ON DELETE CASCADE / SET NULL / SET DEFAULT foreign + # keys inside an open transaction would fire those actions when the old + # table is dropped - transform() should refuse instead + fresh_db.conn.execute("PRAGMA foreign_keys=ON") + fresh_db.executescript(""" + CREATE TABLE authors (id INTEGER PRIMARY KEY, name TEXT); + CREATE TABLE books ( + id INTEGER PRIMARY KEY, + title TEXT, + author_id INTEGER REFERENCES authors(id) ON DELETE {} + ); + """.format(on_delete)) + fresh_db["authors"].insert({"id": 1, "name": "Ursula K. Le Guin"}) + fresh_db["books"].insert({"id": 1, "title": "The Dispossessed", "author_id": 1}) + previous_schema = fresh_db["authors"].schema + with fresh_db.atomic(): + with pytest.raises(TransactionError) as excinfo: + fresh_db["authors"].transform(rename={"name": "author_name"}) + message = str(excinfo.value) + assert "books" in message + assert "ON DELETE {}".format(on_delete.upper()) in message + # Nothing should have changed + assert fresh_db["authors"].schema == previous_schema + assert list(fresh_db["books"].rows) == [ + {"id": 1, "title": "The Dispossessed", "author_id": 1} + ] + assert fresh_db.conn.execute("PRAGMA foreign_keys").fetchone()[0] + + +def test_transform_in_transaction_refuses_self_referential_cascade(fresh_db): + # The copied table carries a foreign key referencing the original table + # name, so a self-referential cascade would wipe the copy too + fresh_db.conn.execute("PRAGMA foreign_keys=ON") + fresh_db.executescript(""" + CREATE TABLE categories ( + id INTEGER PRIMARY KEY, + name TEXT, + parent_id INTEGER REFERENCES categories(id) ON DELETE CASCADE + ); + """) + fresh_db["categories"].insert_all( + [ + {"id": 1, "name": "Fiction", "parent_id": None}, + {"id": 2, "name": "Science Fiction", "parent_id": 1}, + ] + ) + with fresh_db.atomic(): + with pytest.raises(TransactionError) as excinfo: + fresh_db["categories"].transform(rename={"name": "title"}) + assert "categories" in str(excinfo.value) + assert fresh_db["categories"].count == 2 + + +def test_transform_in_transaction_allowed_with_no_action_foreign_key(fresh_db): + # An inbound foreign key without a destructive ON DELETE action is safe + # inside a transaction thanks to PRAGMA defer_foreign_keys + fresh_db.conn.execute("PRAGMA foreign_keys=ON") + fresh_db.executescript(""" + CREATE TABLE authors (id INTEGER PRIMARY KEY, name TEXT); + CREATE TABLE books ( + id INTEGER PRIMARY KEY, + title TEXT, + author_id INTEGER REFERENCES authors(id) + ); + """) + fresh_db["authors"].insert({"id": 1, "name": "Ursula K. Le Guin"}) + fresh_db["books"].insert({"id": 1, "title": "The Dispossessed", "author_id": 1}) + with fresh_db.atomic(): + fresh_db["authors"].transform(rename={"name": "author_name"}) + assert list(fresh_db["authors"].rows) == [ + {"id": 1, "author_name": "Ursula K. Le Guin"} + ] + assert list(fresh_db["books"].rows) == [ + {"id": 1, "title": "The Dispossessed", "author_id": 1} + ] + assert fresh_db.conn.execute("PRAGMA foreign_keys").fetchone()[0] + + +def test_transform_in_transaction_allowed_for_child_table(fresh_db): + # The table being transformed only has an outbound foreign key - dropping + # it fires no ON DELETE actions, so this is allowed inside a transaction + fresh_db.conn.execute("PRAGMA foreign_keys=ON") + fresh_db.executescript(""" + CREATE TABLE authors (id INTEGER PRIMARY KEY, name TEXT); + CREATE TABLE books ( + id INTEGER PRIMARY KEY, + title TEXT, + author_id INTEGER REFERENCES authors(id) ON DELETE CASCADE + ); + """) + fresh_db["authors"].insert({"id": 1, "name": "Ursula K. Le Guin"}) + fresh_db["books"].insert({"id": 1, "title": "The Dispossessed", "author_id": 1}) + with fresh_db.atomic(): + fresh_db["books"].transform(rename={"title": "book_title"}) + assert list(fresh_db["books"].rows) == [ + {"id": 1, "book_title": "The Dispossessed", "author_id": 1} + ] + + +def test_transform_in_transaction_allowed_with_foreign_keys_off(fresh_db): + # With PRAGMA foreign_keys off (the default) no cascades can fire, so + # transform inside a transaction is safe even with a CASCADE schema + fresh_db.executescript(""" + CREATE TABLE authors (id INTEGER PRIMARY KEY, name TEXT); + CREATE TABLE books ( + id INTEGER PRIMARY KEY, + title TEXT, + author_id INTEGER REFERENCES authors(id) ON DELETE CASCADE + ); + """) + fresh_db["authors"].insert({"id": 1, "name": "Ursula K. Le Guin"}) + fresh_db["books"].insert({"id": 1, "title": "The Dispossessed", "author_id": 1}) + with fresh_db.atomic(): + fresh_db["authors"].transform(rename={"name": "author_name"}) + assert list(fresh_db["books"].rows) == [ + {"id": 1, "title": "The Dispossessed", "author_id": 1} + ] + + def test_transform_add_foreign_keys_from_scratch(fresh_db): _add_country_city_continent(fresh_db) fresh_db["places"].insert(_CAVEAU) From 458b3ab5b169eff1f8319c44a7c320c68f54d28b Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sun, 12 Jul 2026 13:52:14 -0700 Subject: [PATCH 3/6] Release 4.1.1 Refs #791, #792, #794, #795 --- docs/changelog.rst | 7 +++++++ pyproject.toml | 2 +- 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 5b9355f..a8006c3 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,13 @@ Changelog =========== +.. _v4_1_1: + +4.1.1 (2026-07-12) +------------------ + +- ``table.transform()`` now raises a ``TransactionError`` if called while a transaction is open with ``PRAGMA foreign_keys`` enabled and the table is referenced by foreign keys with destructive ``ON DELETE`` actions - ``CASCADE``, ``SET NULL`` or ``SET DEFAULT``. The pragma cannot be changed inside a transaction, so previously dropping the old table as part of the transform could fire those actions and silently delete or modify referencing rows. See :ref:`python_api_transform_foreign_keys_transactions` for details and workarounds. (:issue:`794`) +- The CLI and Python API documentation now cross-reference each other: CLI sections link to the equivalent Python API functionality and Python API sections link back to the corresponding CLI command. (:issue:`791`) .. _v4_1: 4.1 (2026-07-11) diff --git a/pyproject.toml b/pyproject.toml index 971f5a0..003322c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "sqlite-utils" -version = "4.1" +version = "4.1.1" description = "CLI tool and Python library for manipulating SQLite databases" readme = { file = "README.md", content-type = "text/markdown" } authors = [ From a947dc673923ff6e95b41d3dfabe1cbd95e6de86 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sun, 12 Jul 2026 17:14:01 -0700 Subject: [PATCH 4/6] Changelog now links to CLI and Python API in most recent entry --- docs/changelog.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index a8006c3..4c868f4 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -10,7 +10,7 @@ ------------------ - ``table.transform()`` now raises a ``TransactionError`` if called while a transaction is open with ``PRAGMA foreign_keys`` enabled and the table is referenced by foreign keys with destructive ``ON DELETE`` actions - ``CASCADE``, ``SET NULL`` or ``SET DEFAULT``. The pragma cannot be changed inside a transaction, so previously dropping the old table as part of the transform could fire those actions and silently delete or modify referencing rows. See :ref:`python_api_transform_foreign_keys_transactions` for details and workarounds. (:issue:`794`) -- The CLI and Python API documentation now cross-reference each other: CLI sections link to the equivalent Python API functionality and Python API sections link back to the corresponding CLI command. (:issue:`791`) +- The :ref:`CLI ` and :ref:`Python API ` documentation now cross-reference each other: CLI sections link to the equivalent Python API functionality and Python API sections link back to the corresponding CLI command. (:issue:`791`) .. _v4_1: 4.1 (2026-07-11) From 69a1c0d960abb20ac03a085142bd59f7fbe002f7 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sat, 25 Jul 2026 14:53:12 -0700 Subject: [PATCH 5/6] Fixes for Ruff>=0.16.0 (#814) * Automated upgrades by Ruff uvx --with 'ruff>=0.16.0' ruff check . --fix --unsafe-fixes * Fix remaining Ruff errors with GPT-5.6 Sol high https://gist.github.com/simonw/6da7906a9fea6e90da131c21a9055199 * Fix flake E501 long lines * New Protocol for migrations to make ty happy --- .gitignore | 1 + docs/cli-reference.rst | 4 +- docs/conf.py | 11 +- pyproject.toml | 9 +- sqlite_utils/__init__.py | 7 +- sqlite_utils/cli.py | 199 +++--- sqlite_utils/db.py | 988 +++++++++++++---------------- sqlite_utils/hookspecs.py | 3 +- sqlite_utils/migrations.py | 27 +- sqlite_utils/plugins.py | 10 +- sqlite_utils/recipes.py | 12 +- sqlite_utils/utils.py | 111 ++-- tests/conftest.py | 5 +- tests/test_analyze_tables.py | 10 +- tests/test_atomic.py | 81 ++- tests/test_cli.py | 88 +-- tests/test_cli_bulk.py | 8 +- tests/test_cli_convert.py | 16 +- tests/test_cli_insert.py | 17 +- tests/test_cli_memory.py | 13 +- tests/test_cli_migrate.py | 3 +- tests/test_column_affinity.py | 3 +- tests/test_constructor.py | 6 +- tests/test_convert.py | 3 +- tests/test_create.py | 68 +- tests/test_create_view.py | 1 + tests/test_default_value.py | 2 +- tests/test_delete.py | 2 +- tests/test_docs.py | 12 +- tests/test_duplicate.py | 6 +- tests/test_enable_counts.py | 10 +- tests/test_extract.py | 15 +- tests/test_extracts.py | 19 +- tests/test_foreign_keys.py | 3 +- tests/test_fts.py | 26 +- tests/test_get.py | 1 + tests/test_gis.py | 5 +- tests/test_hypothesis.py | 3 +- tests/test_insert_files.py | 14 +- tests/test_introspect.py | 7 +- tests/test_list_mode.py | 1 + tests/test_lookup.py | 3 +- tests/test_m2m.py | 6 +- tests/test_migrations.py | 8 +- tests/test_plugins.py | 13 +- tests/test_query.py | 3 +- tests/test_recipes.py | 6 +- tests/test_recreate.py | 6 +- tests/test_rows_from_file.py | 6 +- tests/test_sniff.py | 6 +- tests/test_suggest_column_types.py | 4 +- tests/test_tracer.py | 54 +- tests/test_transform.py | 33 +- tests/test_update.py | 2 +- tests/test_upsert.py | 5 +- tests/test_utils.py | 6 +- tests/test_wal.py | 32 +- 57 files changed, 974 insertions(+), 1049 deletions(-) diff --git a/.gitignore b/.gitignore index 6743708..5b5d2c6 100644 --- a/.gitignore +++ b/.gitignore @@ -15,6 +15,7 @@ venv .schema .vscode .hypothesis +.claude/ Pipfile Pipfile.lock uv.lock diff --git a/docs/cli-reference.rst b/docs/cli-reference.rst index 9fafe28..a4ec402 100644 --- a/docs/cli-reference.rst +++ b/docs/cli-reference.rst @@ -662,7 +662,7 @@ See :ref:`cli_convert`. Convert a string like a,b,c into a JSON array ["a", "b", "c"] r.parsedate(value: 'str', dayfirst: 'bool' = False, yearfirst: 'bool' = False, - errors: 'Optional[object]' = None) -> 'Optional[str]' + errors: 'object | None' = None) -> 'str | None' Parse a date and convert it to ISO date format: yyyy-mm-dd - dayfirst=True: treat xx as the day in xx/yy/zz @@ -671,7 +671,7 @@ See :ref:`cli_convert`. - errors=r.SET_NULL to set values that cannot be parsed to null r.parsedatetime(value: 'str', dayfirst: 'bool' = False, yearfirst: 'bool' = - False, errors: 'Optional[object]' = None) -> 'Optional[str]' + False, errors: 'object | None' = None) -> 'str | None' Parse a datetime and convert it to ISO datetime format: yyyy-mm-ddTHH:MM:SS - dayfirst=True: treat xx as the day in xx/yy/zz diff --git a/docs/conf.py b/docs/conf.py index 4f29b39..62d4642 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -1,10 +1,7 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - import inspect -from pathlib import Path -from subprocess import Popen, PIPE, check_output import sys +from pathlib import Path +from subprocess import PIPE, CalledProcessError, Popen, check_output # This file is execfile()d with the current directory set to its # containing dir. @@ -50,7 +47,7 @@ extlinks = { def _linkcode_git_ref(): try: return check_output(["git", "rev-parse", "HEAD"]).decode("utf8").strip() - except Exception: + except (CalledProcessError, OSError): return "main" @@ -79,7 +76,7 @@ def linkcode_resolve(domain, info): obj = inspect.unwrap(obj) source_file = inspect.getsourcefile(obj) _, line_number = inspect.getsourcelines(obj) - except Exception: + except (OSError, TypeError, ValueError): return None if source_file is None: diff --git a/pyproject.toml b/pyproject.toml index 003322c..6bc0a64 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -79,7 +79,14 @@ build-backend = "setuptools.build_meta" max-line-length = 160 # Black compatibility, E203 whitespace before ':': extend-ignore = ["E203"] -extend-exclude = [".venv", "build", "dist", "docs", "sqlite_utils.egg-info"] +extend-exclude = [ + ".venv", + ".claude", + "build", + "dist", + "docs", + "sqlite_utils.egg-info", +] [tool.setuptools.package-data] sqlite_utils = ["py.typed"] diff --git a/sqlite_utils/__init__.py b/sqlite_utils/__init__.py index 58ee7ab..0d25716 100644 --- a/sqlite_utils/__init__.py +++ b/sqlite_utils/__init__.py @@ -1,7 +1,6 @@ -from .utils import suggest_column_types -from .hookspecs import hookimpl -from .hookspecs import hookspec from .db import Database +from .hookspecs import hookimpl, hookspec from .migrations import Migrations +from .utils import suggest_column_types -__all__ = ["Database", "Migrations", "suggest_column_types", "hookimpl", "hookspec"] +__all__ = ["Database", "Migrations", "hookimpl", "hookspec", "suggest_column_types"] diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index e0b8969..dab4b67 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -1,17 +1,30 @@ import base64 +import csv as csv_std import difflib -from typing import Any -import click -from click_default_group import DefaultGroup -from datetime import datetime, timezone import hashlib +import inspect +import io +import itertools +import json +import os import pathlib +import pdb # noqa: T100 +import sys +import textwrap +from datetime import datetime, timezone from runpy import run_module +from typing import Any + +import click +import tabulate +from click_default_group import DefaultGroup + import sqlite_utils +from sqlite_utils import recipes from sqlite_utils.db import ( + DEFAULT, AlterError, BadMultiValues, - DEFAULT, DescIndex, InvalidColumns, NoTable, @@ -19,36 +32,28 @@ from sqlite_utils.db import ( PrimaryKeyRequired, quote_identifier, ) -from sqlite_utils.plugins import ensure_plugins_loaded, pm, get_plugins +from sqlite_utils.plugins import ensure_plugins_loaded, get_plugins, pm from sqlite_utils.utils import maximize_csv_field_size_limit -from sqlite_utils import recipes -import textwrap -import inspect -import io -import itertools -import json -import os -import pdb -import sys -import csv as csv_std -import tabulate + from .utils import ( + Format, OperationalError, + TypeTracker, _compile_code, chunks, + decode_base64_values, dedupe_keys, file_progress, find_spatialite, - flatten as _flatten, - sqlite3, - decode_base64_values, progressbar, rows_from_file, - Format, - TypeTracker, + sqlite3, +) +from .utils import ( + flatten as _flatten, ) -CONTEXT_SETTINGS = dict(help_option_names=["-h", "--help"]) +CONTEXT_SETTINGS = {"help_option_names": ["-h", "--help"]} def _register_db_for_cleanup(db): @@ -67,7 +72,7 @@ def _close_databases(ctx): for db in ctx.meta.get("_databases_to_close", []): try: db.close() - except Exception: + except sqlite3.Error: pass @@ -174,7 +179,6 @@ def functions_option(fn): @click.version_option() def cli(): "Commands for interacting with a SQLite database" - pass @cli.command() @@ -891,7 +895,7 @@ def enable_counts(path, tables, load_extension): # Check all tables exist bad_tables = [table for table in tables if not db[table].exists()] if bad_tables: - raise click.ClickException("Invalid tables: {}".format(bad_tables)) + raise click.ClickException(f"Invalid tables: {bad_tables}") for table in tables: db.table(table).enable_counts() @@ -1140,9 +1144,7 @@ def insert_upsert_implementation( ) ): raise click.ClickException( - "{}\n\nTry using --alter to add additional columns".format( - e.args[0] - ) + f"{e.args[0]}\n\nTry using --alter to add additional columns" ) # If we can find sql= and parameters= arguments, show those variables = _find_variables(e.__traceback__, ["sql", "parameters"]) @@ -1240,7 +1242,7 @@ def insert_upsert_implementation( 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))] + headers = [f"untitled_{i + 1}" for i in range(len(first_row))] reader = itertools.chain([first_row], reader) else: headers = first_row @@ -1269,9 +1271,7 @@ def insert_upsert_implementation( docs = [docs] except json.decoder.JSONDecodeError as ex: raise click.ClickException( - "Invalid JSON - use --csv for CSV or --tsv for TSV files\n\nJSON error: {}".format( - ex - ) + f"Invalid JSON - use --csv for CSV or --tsv for TSV files\n\nJSON error: {ex}" ) if flatten: docs = (_flatten(doc) for doc in docs) @@ -1290,7 +1290,7 @@ def insert_upsert_implementation( docs = (fn(doc["line"]) for doc in docs) elif text: # Special case: this is allowed to be an iterable - text_value = list(docs)[0]["text"] + text_value = next(iter(docs))["text"] fn_return = fn(text_value) if isinstance(fn_return, dict): docs = [fn_return] @@ -1774,17 +1774,14 @@ def create_table( ctype = columns.pop(0) if ctype.upper() not in VALID_COLUMN_TYPES: raise click.ClickException( - "column types must be one of {}".format(VALID_COLUMN_TYPES) + f"column types must be one of {VALID_COLUMN_TYPES}" ) coltypes[name] = ctype.upper() # Does table already exist? - if table in db.table_names(): - if not ignore and not replace and not transform: - raise click.ClickException( - 'Table "{}" already exists. Use --replace to delete and replace it.'.format( - table - ) - ) + if table in db.table_names() and not ignore and not replace and not transform: + raise click.ClickException( + f'Table "{table}" already exists. Use --replace to delete and replace it.' + ) db.table(table).create( coltypes, pk=pks[0] if len(pks) == 1 else pks, @@ -1819,7 +1816,7 @@ def duplicate(path, table, new_table, ignore, load_extension): db.table(table).duplicate(new_table) except NoTable: if not ignore: - raise click.ClickException('Table "{}" does not exist'.format(table)) + raise click.ClickException(f'Table "{table}" does not exist') @cli.command(name="rename-table") @@ -1843,9 +1840,7 @@ def rename_table(path, table, new_name, ignore, load_extension): db.rename_table(table, new_name) except sqlite3.OperationalError as ex: if not ignore: - raise click.ClickException( - 'Table "{}" could not be renamed. {}'.format(table, str(ex)) - ) + raise click.ClickException(f'Table "{table}" could not be renamed. {ex!s}') @cli.command(name="drop-table") @@ -1874,10 +1869,10 @@ def drop_table(path, table, ignore, load_extension): # A view exists with this name if not ignore: raise click.ClickException( - '"{}" is a view, not a table - use drop-view to drop it'.format(table) + f'"{table}" is a view, not a table - use drop-view to drop it' ) except OperationalError: - raise click.ClickException('Table "{}" does not exist'.format(table)) + raise click.ClickException(f'Table "{table}" does not exist') @cli.command(name="create-view") @@ -1919,9 +1914,7 @@ def create_view(path, view, select, ignore, replace, load_extension): db.view(view).drop() else: raise click.ClickException( - 'View "{}" already exists. Use --replace to delete and replace it.'.format( - view - ) + f'View "{view}" already exists. Use --replace to delete and replace it.' ) db.create_view(view, select) @@ -1953,9 +1946,9 @@ def drop_view(path, view, ignore, load_extension): return if view in db.table_names(): raise click.ClickException( - '"{}" is a table, not a view - use drop-table to drop it'.format(view) + f'"{view}" is a table, not a view - use drop-table to drop it' ) - raise click.ClickException('View "{}" does not exist'.format(view)) + raise click.ClickException(f'View "{view}" does not exist') @cli.command() @@ -2177,7 +2170,7 @@ def memory( file_path = pathlib.Path(path) stem = file_path.stem if stem_counts.get(stem): - file_table = "{}_{}".format(stem, stem_counts[stem]) + file_table = f"{stem}_{stem_counts[stem]}" else: file_table = stem stem_counts[stem] = stem_counts.get(stem, 1) + 1 @@ -2196,14 +2189,14 @@ def memory( if tracker is not None and db.table(file_table).exists(): db.table(file_table).transform(types=tracker.types) # Add convenient t / t1 / t2 views - view_names = ["t{}".format(i + 1)] + view_names = [f"t{i + 1}"] if i == 0: view_names.append("t") for view_name in view_names: if not db[view_name].exists(): db.create_view( view_name, - "select * from {}".format(quote_identifier(file_table)), + f"select * from {quote_identifier(file_table)}", ) finally: if should_close_fp and fp: @@ -2373,19 +2366,17 @@ def search( # Check table exists table_obj = db.table(dbtable) if not table_obj.exists(): - raise click.ClickException("Table '{}' does not exist".format(dbtable)) + raise click.ClickException(f"Table '{dbtable}' does not exist") if not table_obj.detect_fts(): raise click.ClickException( - "Table '{}' is not configured for full-text search".format(dbtable) + f"Table '{dbtable}' is not configured for full-text search" ) if column: # Check they all exist table_columns = table_obj.columns_dict for c in column: if c not in table_columns: - raise click.ClickException( - "Table '{}' has no column '{}".format(dbtable, c) - ) + raise click.ClickException(f"Table '{dbtable}' has no column '{c}") sql = table_obj.search_sql(columns=column, order_by=order, limit=limit) if show_sql: click.echo(sql) @@ -2412,7 +2403,7 @@ def search( except click.ClickException as e: if "malformed MATCH expression" in str(e) or "unterminated string" in str(e): raise click.ClickException( - "{}\n\nTry running this again with the --quote option".format(str(e)) + f"{e!s}\n\nTry running this again with the --quote option" ) else: raise @@ -2479,15 +2470,15 @@ def rows( columns = "*" if column: columns = ", ".join(quote_identifier(c) for c in column) - sql = "select {} from {}".format(columns, quote_identifier(dbtable)) + sql = f"select {columns} from {quote_identifier(dbtable)}" if where: sql += " where " + where if order: sql += " order by " + order if limit: - sql += " limit {}".format(limit) + sql += f" limit {limit}" if offset: - sql += " offset {}".format(offset) + sql += f" offset {offset}" ctx.invoke( query, path=path, @@ -2760,7 +2751,7 @@ def transform( for column, ctype in type: if ctype.upper() not in VALID_COLUMN_TYPES: raise click.ClickException( - "column types must be one of {}".format(VALID_COLUMN_TYPES) + f"column types must be one of {VALID_COLUMN_TYPES}" ) types[column] = ctype.upper() @@ -2858,12 +2849,12 @@ def extract( db = sqlite_utils.Database(path) _register_db_for_cleanup(db) _load_extensions(db, load_extension) - kwargs: dict[str, Any] = dict( - columns=columns, - table=other_table, - fk_column=fk_column, - rename=dict(rename), - ) + kwargs: dict[str, Any] = { + "columns": columns, + "table": other_table, + "fk_column": fk_column, + "rename": dict(rename), + } try: db.table(table).extract(**kwargs) except (NoTable, InvalidColumns) as e: @@ -2958,7 +2949,7 @@ def insert_files( with progressbar(paths_and_relative_paths, silent=silent) as bar: def to_insert(): - for path, relative_path in bar: + for file_path, relative_path in bar: row = {} # content_text is special case as it considers 'encoding' @@ -2970,19 +2961,21 @@ def insert_files( raise UnicodeDecodeErrorForPath(e, resolved) lookups = dict(FILE_COLUMNS, content_text=_content_text) - if path == "-": + if file_path == "-": stdin_data = sys.stdin.buffer.read() # We only support a subset of columns for this case lookups = { "name": lambda p: name or "-", "path": lambda p: name or "-", - "content": lambda p: stdin_data, - "content_text": lambda p: stdin_data.decode( + "content": lambda p, data=stdin_data: data, + "content_text": lambda p, data=stdin_data: data.decode( encoding or "utf-8" ), - "sha256": lambda p: hashlib.sha256(stdin_data).hexdigest(), - "md5": lambda p: hashlib.md5(stdin_data).hexdigest(), - "size": lambda p: len(stdin_data), + "sha256": lambda p, data=stdin_data: hashlib.sha256( + data + ).hexdigest(), + "md5": lambda p, data=stdin_data: hashlib.md5(data).hexdigest(), + "size": lambda p, data=stdin_data: len(data), } for coldef in column: if ":" in coldef: @@ -2990,7 +2983,7 @@ def insert_files( else: colname, coltype = coldef, coldef try: - value = lookups[coltype](path) + value = lookups[coltype](file_path) row[colname] = value except KeyError: raise click.ClickException( @@ -3018,7 +3011,7 @@ def insert_files( except UnicodeDecodeErrorForPath as e: raise click.ClickException( UNICODE_ERROR.format( - "Could not read file '{}' as text\n\n{}".format(e.path, e.exception) + f"Could not read file '{e.path}' as text\n\n{e.exception}" ) ) @@ -3196,7 +3189,7 @@ def _generate_convert_help(): for name in recipe_names: fn = getattr(recipes, name) doc = textwrap.dedent(fn.__doc__.rstrip()).replace("\b\n", "") - help += "\n\nr.{}{}\n\n\b{}".format(name, str(inspect.signature(fn)), doc) + help += f"\n\nr.{name}{inspect.signature(fn)!s}\n\n\b{doc}" help += "\n\n" help += textwrap.dedent(""" You can use these recipes like so: @@ -3299,7 +3292,7 @@ def convert( """.format( column=columns[0], table=table, - where=" where {}".format(where) if where is not None else "", + where=f" where {where}" if where is not None else "", ) for row in db.conn.execute(sql, where_args).fetchall(): click.echo(str(row[0])) @@ -3319,7 +3312,7 @@ def convert( def wrapped_fn(value): try: return fn_(value) - except Exception as ex: + except Exception as ex: # noqa: BLE001 print("\nException raised, dropping into pdb...:", ex) pdb.post_mortem(ex.__traceback__) sys.exit(1) @@ -3339,9 +3332,7 @@ def convert( ) except BadMultiValues as e: raise click.ClickException( - "When using --multi code must return a Python dictionary - returned: {}".format( - repr(e.values) - ) + f"When using --multi code must return a Python dictionary - returned: {e.values!r}" ) @@ -3459,7 +3450,7 @@ def create_spatial_index(db_path, table, column_name, load_extension): def _find_migration_files(migrations): if not migrations: - migrations = [pathlib.Path(".").resolve()] + migrations = [pathlib.Path.cwd()] files = set() for path_str in migrations: path = pathlib.Path(path_str) @@ -3484,7 +3475,7 @@ def _load_migration_sets(files): "__file__": str(filepath), "__name__": "__sqlite_utils_migration__", } - exec(code, namespace) + exec(code, namespace) # noqa: S102 migration_sets.extend( obj for obj in namespace.values() if _compatible_migration_set(obj) ) @@ -3493,17 +3484,17 @@ def _load_migration_sets(files): def _display_migration_list(db, migration_sets): for migration_set in migration_sets: - click.echo("Migrations for: {}".format(migration_set.name)) + click.echo(f"Migrations for: {migration_set.name}") click.echo() click.echo(" Applied:") for migration in migration_set.applied(db): - click.echo(" {} - {}".format(migration.name, migration.applied_at)) + click.echo(f" {migration.name} - {migration.applied_at}") click.echo() click.echo(" Pending:") output = False for migration in migration_set.pending(db): output = True - click.echo(" {}".format(migration.name)) + click.echo(f" {migration.name}") if not output: click.echo(" (none)") click.echo() @@ -3583,7 +3574,7 @@ def migrate(db_path, migrations, stop_before, list_, verbose): prev_schema = db.schema if verbose: - click.echo("Migrating {}".format(db_path)) + click.echo(f"Migrating {db_path}") click.echo("\nSchema before:\n") click.echo(textwrap.indent(prev_schema, " ") or " (empty)") click.echo() @@ -3594,9 +3585,7 @@ def migrate(db_path, migrations, stop_before, list_, verbose): names = {m.name for m in migration_set.pending(db)} names.update(m.name for m in migration_set.applied(db)) known_names.update(names) - known_names.update( - "{}:{}".format(migration_set.name, name) for name in names - ) + known_names.update(f"{migration_set.name}:{name}" for name in names) unknown = [value for value in stop_before if value not in known_names] if unknown: raise click.ClickException( @@ -3652,7 +3641,7 @@ def _render_common(title, values): return "" lines = [title] for value, count in values: - lines.append(" {}: {}".format(count, value)) + lines.append(f" {count}: {value}") return "\n".join(lines) @@ -3722,7 +3711,7 @@ def maybe_json(value): if not isinstance(value, str): return value stripped = value.strip() - if not (stripped.startswith("{") or stripped.startswith("[")): + if not (stripped.startswith(("{", "["))): return value try: return json.loads(stripped) @@ -3740,7 +3729,7 @@ def json_binary(value): def verify_is_dict(doc): if not isinstance(doc, dict): raise click.ClickException( - "Rows must all be dictionaries, got: {}".format(repr(doc)[:1000]) + f"Rows must all be dictionaries, got: {repr(doc)[:1000]}" ) return doc @@ -3768,14 +3757,14 @@ def _register_functions(db, functions): try: functions = pathlib.Path(functions).read_text() except FileNotFoundError: - raise click.ClickException("File not found: {}".format(functions)) + raise click.ClickException(f"File not found: {functions}") sqlite3.enable_callback_tracebacks(True) globals = {} try: - exec(functions, globals) + exec(functions, globals) # noqa: S102 except SyntaxError as ex: - raise click.ClickException("Error in functions definition: {}".format(ex)) + raise click.ClickException(f"Error in functions definition: {ex}") # Register all callables in the locals dict: for name, value in globals.items(): if callable(value) and not name.startswith("_"): @@ -3796,12 +3785,12 @@ def _rows_from_code(code): try: code = pathlib.Path(code).read_text() except FileNotFoundError: - raise click.ClickException("File not found: {}".format(code)) + raise click.ClickException(f"File not found: {code}") namespace = {} try: - exec(code, namespace) + exec(code, namespace) # noqa: S102 except SyntaxError as ex: - raise click.ClickException("Error in --code: {}".format(ex)) + raise click.ClickException(f"Error in --code: {ex}") rows = namespace.get("rows") if callable(rows): rows = rows() diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index e97b7d9..9a00123 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -1,19 +1,4 @@ -from .utils import ( - chunks, - dedupe_keys, - hash_record, - sqlite3, - OperationalError, - suggest_column_types, - types_for_column_types, - column_affinity, - progressbar, - find_spatialite, -) import binascii -from collections import namedtuple -from dataclasses import dataclass, field -from collections.abc import Mapping import contextlib import datetime import decimal @@ -25,26 +10,36 @@ import os import pathlib import re import secrets -from sqlite_fts4 import rank_bm25 import textwrap -from typing import ( - cast, - Any, - Callable, - Dict, - Generator, - Iterable, - Sequence, - Set, - Type, - Union, - Optional, - List, - Tuple, -) import uuid +from collections import namedtuple +from collections.abc import Callable, Generator, Iterable, Mapping, Sequence +from dataclasses import dataclass, field +from types import TracebackType +from typing import ( + Any, + Union, + cast, +) + +from sqlite_fts4 import rank_bm25 +from typing_extensions import Self + from sqlite_utils.plugins import ensure_plugins_loaded, pm +from .utils import ( + OperationalError, + chunks, + column_affinity, + dedupe_keys, + find_spatialite, + hash_record, + progressbar, + sqlite3, + suggest_column_types, + types_for_column_types, +) + try: iterdump = importlib.import_module("sqlite_dump").iterdump except ImportError: @@ -226,11 +221,11 @@ class ForeignKey: table: str # column/other_column are None for compound keys, which would break # ordering against str values - comparison uses columns/other_columns - column: Optional[str] = field(compare=False) + column: str | None = field(compare=False) other_table: str - other_column: Optional[str] = field(compare=False) - columns: Tuple[str, ...] = () - other_columns: Tuple[str, ...] = () + other_column: str | None = field(compare=False) + columns: tuple[str, ...] = () + other_columns: tuple[str, ...] = () is_compound: bool = False on_delete: str = "NO ACTION" on_update: str = "NO ACTION" @@ -259,9 +254,9 @@ def _fk_actions_sql(fk: ForeignKey) -> str: "ON UPDATE/ON DELETE clauses for a foreign key, or an empty string." actions = "" if fk.on_update and fk.on_update != "NO ACTION": - actions += " ON UPDATE {}".format(fk.on_update) + actions += f" ON UPDATE {fk.on_update}" if fk.on_delete and fk.on_delete != "NO ACTION": - actions += " ON DELETE {}".format(fk.on_delete) + actions += f" ON DELETE {fk.on_delete}" return actions @@ -278,20 +273,20 @@ class TransformError(Exception): # A single column name, or a tuple of columns for a compound foreign key -ForeignKeyColumns = Union[str, Tuple[str, ...], List[str]] +ForeignKeyColumns = str | tuple[str, ...] | list[str] # (table, column(s), other_table, other_column(s)) -ForeignKeyTuple = Tuple[str, ForeignKeyColumns, str, ForeignKeyColumns] +ForeignKeyTuple = tuple[str, ForeignKeyColumns, str, ForeignKeyColumns] -ForeignKeyIndicator = Union[ - str, - ForeignKey, - Tuple[ForeignKeyColumns, str], - Tuple[ForeignKeyColumns, str, ForeignKeyColumns], - ForeignKeyTuple, -] +ForeignKeyIndicator = ( + str + | ForeignKey + | tuple[ForeignKeyColumns, str] + | tuple[ForeignKeyColumns, str, ForeignKeyColumns] + | ForeignKeyTuple +) -ForeignKeysType = Union[Iterable[ForeignKeyIndicator], List[ForeignKeyIndicator]] +ForeignKeysType = Iterable[ForeignKeyIndicator] | list[ForeignKeyIndicator] class Default: @@ -300,7 +295,7 @@ class Default: DEFAULT = Default() -Tracer = Callable[[str, Optional[Union[Sequence[Any], Dict[str, Any]]]], None] +Tracer = Callable[[str, Sequence[Any] | dict[str, Any] | None], None] def _iter_complete_sql_statements(sql: str) -> Generator[str, None, None]: @@ -316,7 +311,7 @@ def _iter_complete_sql_statements(sql: str) -> Generator[str, None, None]: yield statement_sql -COLUMN_TYPE_MAPPING: Dict[Any, str] = { +COLUMN_TYPE_MAPPING: dict[Any, str] = { float: "REAL", int: "INTEGER", bool: "INTEGER", @@ -512,12 +507,12 @@ class Database: def __init__( self, - filename_or_conn: Optional[Union[str, pathlib.Path, sqlite3.Connection]] = None, + filename_or_conn: str | pathlib.Path | sqlite3.Connection | None = None, memory: bool = False, - memory_name: Optional[str] = None, + memory_name: str | None = None, recreate: bool = False, recursive_triggers: bool = True, - tracer: Optional[Tracer] = None, + tracer: Tracer | None = None, use_counts_table: bool = False, execute_plugins: bool = True, use_old_upsert: bool = False, @@ -532,7 +527,7 @@ class Database: ): raise ValueError("Either specify a filename_or_conn or pass memory=True") if memory_name: - uri = "file:{}?mode=memory&cache=shared".format(memory_name) + uri = f"file:{memory_name}?mode=memory&cache=shared" self.conn = sqlite3.connect( uri, uri=True, @@ -569,7 +564,7 @@ class Database: "transaction handling - connections created with " "autocommit=True or autocommit=False are not supported" ) - self._tracer: Optional[Tracer] = tracer + self._tracer: Tracer | None = tracer if recursive_triggers: self.execute("PRAGMA recursive_triggers=on;") self._registered_functions: set = set() @@ -579,14 +574,14 @@ class Database: pm.hook.prepare_connection(conn=self.conn) self.strict = strict - def __enter__(self) -> "Database": + def __enter__(self) -> Self: return self def __exit__( self, - exc_type: Optional[Type[BaseException]], - exc_val: Optional[BaseException], - exc_tb: Optional[object], + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, ) -> None: self.close() @@ -602,8 +597,8 @@ class Database: Nested blocks use SQLite savepoints. """ if self.conn.in_transaction: - savepoint = "sqlite_utils_{}".format(secrets.token_hex(16)) - self.conn.execute("SAVEPOINT {};".format(savepoint)) + savepoint = f"sqlite_utils_{secrets.token_hex(16)}" + self.conn.execute(f"SAVEPOINT {savepoint};") try: yield self except BaseException: @@ -612,11 +607,11 @@ class Database: # anyway would mask the original exception with # "no such savepoint" if self.conn.in_transaction: - self.conn.execute("ROLLBACK TO SAVEPOINT {};".format(savepoint)) - self.conn.execute("RELEASE SAVEPOINT {};".format(savepoint)) + self.conn.execute(f"ROLLBACK TO SAVEPOINT {savepoint};") + self.conn.execute(f"RELEASE SAVEPOINT {savepoint};") raise else: - self.conn.execute("RELEASE SAVEPOINT {};".format(savepoint)) + self.conn.execute(f"RELEASE SAVEPOINT {savepoint};") else: self.conn.execute("BEGIN") try: @@ -695,9 +690,7 @@ class Database: self.conn.isolation_level = old_isolation_level @contextlib.contextmanager - def tracer( - self, tracer: Optional[Tracer] = None - ) -> Generator["Database", None, None]: + def tracer(self, tracer: Tracer | None = None) -> Generator["Database", None, None]: """ Context manager to temporarily set a tracer function - all executed SQL queries will be passed to this. @@ -734,15 +727,15 @@ class Database: return self.table(table_name) def __repr__(self) -> str: - return "".format(self.conn) + return f"" def register_function( self, - fn: Optional[Callable] = None, + fn: Callable | None = None, deterministic: bool = False, replace: bool = False, - name: Optional[str] = None, - ) -> Optional[Callable[[Callable], Callable]]: + name: str | None = None, + ) -> Callable[[Callable], Callable] | None: """ ``fn`` will be made available as a function within SQL, with the same name and number of arguments. Can be used as a decorator:: @@ -770,7 +763,7 @@ class Database: arity = len(inspect.signature(fn).parameters) if not replace and (fn_name, arity) in self._registered_functions: return fn - kwargs: Dict[str, bool] = {} + kwargs: dict[str, bool] = {} registered = False if deterministic: # Try this, but fall back if sqlite3.NotSupportedError @@ -796,7 +789,7 @@ class Database: "Register the ``rank_bm25(match_info)`` function used for calculating relevance with SQLite FTS4." self.register_function(rank_bm25, deterministic=True, replace=True) - def attach(self, alias: str, filepath: Union[str, pathlib.Path]) -> None: + def attach(self, alias: str, filepath: str | pathlib.Path) -> None: """ Attach another SQLite database file to this connection with the specified alias, equivalent to:: @@ -805,15 +798,13 @@ class Database: :param alias: Alias name to use :param filepath: Path to SQLite database file on disk """ - attach_sql = """ - ATTACH DATABASE '{}' AS {}; - """.format( - str(pathlib.Path(filepath).resolve()), quote_identifier(alias) - ).strip() + attach_sql = f""" + ATTACH DATABASE '{pathlib.Path(filepath).resolve()!s}' AS {quote_identifier(alias)}; + """.strip() self.execute(attach_sql) def query( - self, sql: str, params: Optional[Union[Sequence, Dict[str, Any]]] = None + self, sql: str, params: Sequence | dict[str, Any] | None = None ) -> Generator[dict, None, None]: """ Execute ``sql`` and return an iterable of dictionaries representing each row. @@ -891,7 +882,7 @@ class Database: self.conn.execute('RELEASE "sqlite_utils_query"') def execute( - self, sql: str, parameters: Optional[Union[Sequence, Dict[str, Any]]] = None + self, sql: str, parameters: Sequence | dict[str, Any] | None = None ) -> sqlite3.Cursor: """ Execute SQL query and return a ``sqlite3.Cursor``. @@ -960,7 +951,7 @@ class Database: :param table_name: Name of the table """ if table_name in self.view_names(): - raise NoTable("Table {} is actually a view".format(table_name)) + raise NoTable(f"Table {table_name} is actually a view") kwargs.setdefault("strict", self.strict) return Table(self, table_name, **kwargs) @@ -973,11 +964,9 @@ class Database: if view_name not in self.view_names(): if view_name in self.table_names(): raise NoView( - "View {name} does not exist - {name} is a table".format( - name=view_name - ) + f"View {view_name} does not exist - {view_name} is a table" ) - raise NoView("View {} does not exist".format(view_name)) + raise NoView(f"View {view_name} does not exist") return View(self, view_name) def quote(self, value: str) -> str: @@ -1013,9 +1002,7 @@ class Database: query += '"' bits = _quote_fts_re.split(query) bits = [b for b in bits if b and b != '""'] - return " ".join( - '"{}"'.format(bit) if not bit.startswith('"') else bit for bit in bits - ) + return " ".join(f'"{bit}"' if not bit.startswith('"') else bit for bit in bits) def quote_default_value(self, value: str) -> str: if any( @@ -1036,11 +1023,11 @@ class Database: if str(value).endswith(")"): # Expr - return "({})".format(value) + return f"({value})" return self.quote(value) - def table_names(self, fts4: bool = False, fts5: bool = False) -> List[str]: + def table_names(self, fts4: bool = False, fts5: bool = False) -> list[str]: """ List of string table names in this database. @@ -1055,7 +1042,7 @@ class Database: sql = "select name from sqlite_master where {}".format(" AND ".join(where)) return [r[0] for r in self.execute(sql).fetchall()] - def view_names(self) -> List[str]: + def view_names(self) -> list[str]: "List of string view names in this database." return [ r[0] @@ -1065,17 +1052,17 @@ class Database: ] @property - def tables(self) -> List["Table"]: + def tables(self) -> list["Table"]: "List of Table objects in this database." return [self.table(name) for name in self.table_names()] @property - def views(self) -> List["View"]: + def views(self) -> list["View"]: "List of View objects in this database." return [self.view(name) for name in self.view_names()] @property - def triggers(self) -> List[Trigger]: + def triggers(self) -> list[Trigger]: "List of ``(name, table_name, sql)`` tuples representing triggers in this database." return [ Trigger(*r) @@ -1085,7 +1072,7 @@ class Database: ] @property - def triggers_dict(self) -> Dict[str, str]: + def triggers_dict(self) -> dict[str, str]: "A ``{trigger_name: sql}`` dictionary of triggers in this database." return {trigger.name: trigger.sql for trigger in self.triggers} @@ -1107,14 +1094,12 @@ class Database: "Does this database support STRICT mode?" if not hasattr(self, "_supports_strict"): try: - table_name = "t{}".format(secrets.token_hex(16)) + table_name = f"t{secrets.token_hex(16)}" with self.atomic(): - self.conn.execute( - "create table {} (name text) strict".format(table_name) - ) - self.conn.execute("drop table {}".format(table_name)) + self.conn.execute(f"create table {table_name} (name text) strict") + self.conn.execute(f"drop table {table_name}") self._supports_strict = True - except Exception: + except sqlite3.OperationalError: self._supports_strict = False return self._supports_strict @@ -1122,32 +1107,28 @@ class Database: def supports_on_conflict(self) -> bool: # SQLite's upsert is implemented as INSERT INTO ... ON CONFLICT DO ... if not hasattr(self, "_supports_on_conflict"): - table_name = "t{}".format(secrets.token_hex(16)) + table_name = f"t{secrets.token_hex(16)}" try: with self.atomic(): self.conn.execute( - "create table {} (id integer primary key, name text)".format( - table_name - ) + f"create table {table_name} (id integer primary key, name text)" ) self.conn.execute( - "insert into {} (id, name) values (1, 'one')".format(table_name) + f"insert into {table_name} (id, name) values (1, 'one')" ) self.conn.execute( - ( - "insert into {} (id, name) values (1, 'two') " - "on conflict do update set name = 'two'" - ).format(table_name) + f"insert into {table_name} (id, name) values (1, 'two') " + "on conflict do update set name = 'two'" ) self._supports_on_conflict = True - except Exception: + except sqlite3.OperationalError: self._supports_on_conflict = False finally: - self.conn.execute("drop table if exists {}".format(table_name)) + self.conn.execute(f"drop table if exists {table_name}") return self._supports_on_conflict @property - def sqlite_version(self) -> Tuple[int, ...]: + def sqlite_version(self) -> tuple[int, ...]: "Version of SQLite, as a tuple of integers for example ``(3, 36, 0)``." row = self.execute("select sqlite_version()").fetchall()[0] return tuple(map(int, row[0].split("."))) @@ -1191,7 +1172,7 @@ class Database: # guarantee of atomic() and of user-managed transactions if self.conn.in_transaction: raise TransactionError( - "{} cannot be used while a transaction is open".format(operation) + f"{operation} cannot be used while a transaction is open" ) def _ensure_counts_table(self) -> None: @@ -1212,14 +1193,14 @@ class Database: table.enable_counts() self.use_counts_table = True - def cached_counts(self, tables: Optional[Iterable[str]] = None) -> Dict[str, int]: + def cached_counts(self, tables: Iterable[str] | None = None) -> dict[str, int]: """ Return ``{table_name: count}`` dictionary of cached counts for specified tables, or all tables if ``tables`` not provided. :param tables: Subset list of tables to return counts for. """ - sql = 'select "table", count from {}'.format(self._counts_table_name) + sql = f'select "table", count from {self._counts_table_name}' tables_list = list(tables) if tables else None if tables_list: sql += ' where "table" in ({})'.format(", ".join("?" for _ in tables_list)) @@ -1241,13 +1222,13 @@ class Database: ) def execute_returning_dicts( - self, sql: str, params: Optional[Union[Sequence, Dict[str, Any]]] = None - ) -> List[dict]: + self, sql: str, params: Sequence | dict[str, Any] | None = None + ) -> list[dict]: return list(self.query(sql, params)) def resolve_foreign_keys( self, name: str, foreign_keys: ForeignKeysType - ) -> List[ForeignKey]: + ) -> list[ForeignKey]: """ Given a list of differing foreign_keys definitions, return a list of fully resolved ForeignKey() named tuples. @@ -1274,7 +1255,7 @@ class Database: fks.append(ForeignKey(name, fk, other_table, other_column)) continue if not isinstance(fk, (tuple, list)): - raise ValueError( + raise ValueError( # noqa: TRY004 "foreign_keys= should be a list of tuples, " "ForeignKey objects or column name strings" ) @@ -1282,9 +1263,7 @@ class Database: if len(tuple_or_list) == 4: if tuple_or_list[0] != name: raise ValueError( - "First item in {} should have been {}".format( - tuple_or_list, name - ) + f"First item in {tuple_or_list} should have been {name}" ) tuple_or_list = tuple_or_list[1:] if len(tuple_or_list) not in (2, 3): @@ -1299,8 +1278,8 @@ class Database: if len(tuple_or_list) == 3: if not isinstance(tuple_or_list[2], (list, tuple)): raise ValueError( - "Compound foreign key {} should reference a tuple " - "of other columns".format(tuple(tuple_or_list)) + f"Compound foreign key {tuple(tuple_or_list)} should reference a tuple " + "of other columns" ) other_columns = tuple(tuple_or_list[2]) else: @@ -1308,8 +1287,8 @@ class Database: other_columns = tuple(self.table(other_table).pks) if len(columns) != len(other_columns): raise ValueError( - "Compound foreign key {} should have the same number " - "of columns on both sides".format(tuple(tuple_or_list)) + f"Compound foreign key {tuple(tuple_or_list)} should have the same number " + "of columns on both sides" ) if len(columns) == 1: # Single-column key passed as a one-item list @@ -1389,15 +1368,15 @@ class Database: def create_table_sql( self, name: str, - columns: Dict[str, Any], - pk: Optional[Any] = None, - foreign_keys: Optional[ForeignKeysType] = None, - column_order: Optional[List[str]] = None, - not_null: Optional[Iterable[str]] = None, - defaults: Optional[Dict[str, Any]] = None, - hash_id: Optional[str] = None, - hash_id_columns: Optional[Iterable[str]] = None, - extracts: Optional[Union[Dict[str, str], List[str]]] = None, + columns: dict[str, Any], + pk: Any | None = None, + foreign_keys: ForeignKeysType | None = None, + column_order: list[str] | None = None, + not_null: Iterable[str] | None = None, + defaults: dict[str, Any] | None = None, + hash_id: str | None = None, + hash_id_columns: Iterable[str] | None = None, + extracts: dict[str, str] | list[str] | None = None, if_not_exists: bool = False, strict: bool = False, ) -> str: @@ -1419,7 +1398,7 @@ class Database: """ if hash_id_columns and (hash_id is None): hash_id = "id" - resolved_fks: List[ForeignKey] = [ + resolved_fks: list[ForeignKey] = [ self._resolve_foreign_key_casing(fk, columns) for fk in self.resolve_foreign_keys(name, foreign_keys or []) ] @@ -1449,15 +1428,11 @@ class Database: raise ValueError("Tables must have at least one column") if not all(n in columns for n in not_null): raise ValueError( - "not_null set {} includes items not in columns {}".format( - repr(not_null), repr(set(columns.keys())) - ) + f"not_null set {not_null!r} includes items not in columns {set(columns.keys())!r}" ) if not all(n in columns for n in defaults): raise ValueError( - "defaults set {} includes items not in columns {}".format( - repr(set(defaults)), repr(set(columns.keys())) - ) + f"defaults set {set(defaults)!r} includes items not in columns {set(columns.keys())!r}" ) column_items = list(columns.items()) if column_order is not None: @@ -1477,9 +1452,7 @@ class Database: if other_column != "rowid" and not any( c for c in self[fk.other_table].columns if c.name == other_column ): - raise AlterError( - "No such column: {}.{}".format(fk.other_table, other_column) - ) + raise AlterError(f"No such column: {fk.other_table}.{other_column}") column_defs = [] # ensure pk is a tuple @@ -1500,16 +1473,12 @@ class Database: column_extras.append("NOT NULL") if column_name in defaults and defaults[column_name] is not None: column_extras.append( - "DEFAULT {}".format(self.quote_default_value(defaults[column_name])) + f"DEFAULT {self.quote_default_value(defaults[column_name])}" ) if column_name in foreign_keys_by_column: fk = foreign_keys_by_column[column_name] column_extras.append( - "REFERENCES {}({}){}".format( - quote_identifier(fk.other_table), - quote_identifier(cast(str, fk.other_column)), - _fk_actions_sql(fk), - ) + f"REFERENCES {quote_identifier(fk.other_table)}({quote_identifier(cast(str, fk.other_column))}){_fk_actions_sql(fk)}" ) column_type_str = COLUMN_TYPE_MAPPING[column_type] # Special case for strict tables to map FLOAT to REAL @@ -1566,15 +1535,15 @@ class Database: def create_table( self, name: str, - columns: Dict[str, Any], - pk: Optional[Any] = None, - foreign_keys: Optional[ForeignKeysType] = None, - column_order: Optional[List[str]] = None, - not_null: Optional[Iterable[str]] = None, - defaults: Optional[Dict[str, Any]] = None, - hash_id: Optional[str] = None, - hash_id_columns: Optional[Iterable[str]] = None, - extracts: Optional[Union[Dict[str, str], List[str]]] = None, + columns: dict[str, Any], + pk: Any | None = None, + foreign_keys: ForeignKeysType | None = None, + column_order: list[str] | None = None, + not_null: Iterable[str] | None = None, + defaults: dict[str, Any] | None = None, + hash_id: str | None = None, + hash_id_columns: Iterable[str] | None = None, + extracts: dict[str, str] | list[str] | None = None, if_not_exists: bool = False, replace: bool = False, ignore: bool = False, @@ -1618,11 +1587,11 @@ class Database: resolve_casing(col_name, existing_columns): col_type for col_name, col_type in columns.items() } - missing_columns = dict( - (col_name, col_type) + missing_columns = { + col_name: col_type for col_name, col_type in columns.items() if col_name not in existing_columns - ) + } columns_to_drop = [ column for column in existing_columns if column not in columns ] @@ -1709,9 +1678,7 @@ class Database: :param new_name: Name to rename it to """ self.execute( - "ALTER TABLE {} RENAME TO {}".format( - quote_identifier(name), quote_identifier(new_name) - ) + f"ALTER TABLE {quote_identifier(name)} RENAME TO {quote_identifier(new_name)}" ) def create_view( @@ -1727,23 +1694,20 @@ class Database: """ if ignore and replace: raise ValueError("Use one or the other of ignore/replace, not both") - create_sql = "CREATE VIEW {name} AS {sql}".format( - name=quote_identifier(name), sql=sql - ) - if ignore or replace: - # Does view exist already? - if name in self.view_names(): - if ignore: + create_sql = f"CREATE VIEW {quote_identifier(name)} AS {sql}" + if (ignore or replace) and name in self.view_names(): + # View exists already + if ignore: + return self + elif replace: + # If SQL is the same, do nothing + if create_sql == self[name].schema: return self - elif replace: - # If SQL is the same, do nothing - if create_sql == self[name].schema: - return self - self[name].drop() + self[name].drop() self.execute(create_sql) return self - def m2m_table_candidates(self, table: str, other_table: str) -> List[str]: + def m2m_table_candidates(self, table: str, other_table: str) -> list[str]: """ Given two table names returns the name of tables that could define a many-to-many relationship between those two tables, based on having @@ -1762,7 +1726,7 @@ class Database: return candidates def add_foreign_keys( - self, foreign_keys: Iterable[Union[ForeignKey, ForeignKeyTuple]] + self, foreign_keys: Iterable[ForeignKey | ForeignKeyTuple] ) -> None: """ See :ref:`python_api_add_foreign_keys`. @@ -1782,7 +1746,7 @@ class Database: "(table, column, other_table, other_column)" ) - foreign_keys_to_create: List[ForeignKey] = [] + foreign_keys_to_create: list[ForeignKey] = [] # Verify that all tables and columns exist for fk in foreign_keys: @@ -1823,7 +1787,7 @@ class Database: table = fk_object.table other_table = fk_object.other_table if not self.table(table).exists(): - raise AlterError("No such table: {}".format(table)) + raise AlterError(f"No such table: {table}") table_obj = self.table(table) fk_object = self._resolve_foreign_key_casing( fk_object, table_obj.columns_dict @@ -1832,18 +1796,16 @@ class Database: other_columns = fk_object.other_columns for column in columns: if column not in table_obj.columns_dict: - raise AlterError("No such column: {} in {}".format(column, table)) + raise AlterError(f"No such column: {column} in {table}") if not self[other_table].exists(): - raise AlterError("No such other_table: {}".format(other_table)) + raise AlterError(f"No such other_table: {other_table}") for other_column in other_columns: if ( other_column != "rowid" and other_column not in self[other_table].columns_dict ): raise AlterError( - "No such other_column: {} in {}".format( - other_column, other_table - ) + f"No such other_column: {other_column} in {other_table}" ) # Silently skip foreign keys that exist already - but only if # they match exactly, including ON DELETE/ON UPDATE actions @@ -1874,7 +1836,7 @@ class Database: ) # Group them by table - by_table: Dict[str, List[ForeignKey]] = {} + by_table: dict[str, list[ForeignKey]] = {} for fk_object in foreign_keys_to_create: by_table.setdefault(fk_object.table, []).append(fk_object) @@ -1899,7 +1861,7 @@ class Database: "Run a SQLite ``VACUUM`` against the database." self.execute("VACUUM;") - def analyze(self, name: Optional[str] = None) -> None: + def analyze(self, name: str | None = None) -> None: """ Run ``ANALYZE`` against the entire database or a named table or index. @@ -1907,7 +1869,7 @@ class Database: """ sql = "ANALYZE" if name is not None: - sql += " {}".format(quote_identifier(name)) + sql += f" {quote_identifier(name)}" self.execute(sql) def iterdump(self) -> Generator[str, None, None]: @@ -1922,7 +1884,7 @@ class Database: "conn.iterdump() not found - try pip install sqlite-dump" ) - def init_spatialite(self, path: Optional[str] = None) -> bool: + def init_spatialite(self, path: str | None = None) -> bool: """ The ``init_spatialite`` method will load and initialize the SpatiaLite extension. The ``path`` argument should be an absolute path to the compiled extension, which @@ -1980,8 +1942,8 @@ class Queryable: def count_where( self, - where: Optional[str] = None, - where_args: Optional[Union[Sequence, Dict[str, Any]]] = None, + where: str | None = None, + where_args: Sequence | dict[str, Any] | None = None, ) -> int: """ Executes ``SELECT count(*) FROM table WHERE ...`` and returns a count. @@ -1990,7 +1952,7 @@ class Queryable: :param where_args: Parameters to use with that fragment - an iterable for ``id > ?`` parameters, or a dictionary for ``id > :id`` """ - sql = "select count(*) from {}".format(quote_identifier(self.name)) + sql = f"select count(*) from {quote_identifier(self.name)}" if where is not None: sql += " where " + where return self.db.execute(sql, where_args or []).fetchone()[0] @@ -2005,19 +1967,19 @@ class Queryable: return self.count_where() @property - def rows(self) -> Generator[Dict[str, Any], None, None]: + def rows(self) -> Generator[dict[str, Any], None, None]: "Iterate over every dictionaries for each row in this table or view." return self.rows_where() def rows_where( self, - where: Optional[str] = None, - where_args: Optional[Union[Sequence, Dict[str, Any]]] = None, - order_by: Optional[str] = None, + where: str | None = None, + where_args: Sequence | dict[str, Any] | None = None, + order_by: str | None = None, select: str = "*", - limit: Optional[int] = None, - offset: Optional[int] = None, - ) -> Generator[Dict[str, Any], None, None]: + limit: int | None = None, + offset: int | None = None, + ) -> Generator[dict[str, Any], None, None]: """ Iterate over every row in this table or view that matches the specified where clause. @@ -2033,15 +1995,15 @@ class Queryable: """ if not self.exists(): return - sql = "select {} from {}".format(select, quote_identifier(self.name)) + sql = f"select {select} from {quote_identifier(self.name)}" if where is not None: sql += " where " + where if order_by is not None: sql += " order by " + order_by if limit is not None: - sql += " limit {}".format(limit) + sql += f" limit {limit}" if offset is not None: - sql += " offset {}".format(offset) + sql += f" offset {offset}" cursor = self.db.execute(sql, where_args or []) columns = dedupe_keys(c[0] for c in cursor.description) for row in cursor: @@ -2049,12 +2011,12 @@ class Queryable: def pks_and_rows_where( self, - where: Optional[str] = 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[str, Any]], None, None]: + where: str | None = None, + where_args: Sequence | dict[str, Any] | None = None, + order_by: str | None = None, + limit: int | None = None, + offset: int | 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. @@ -2096,17 +2058,17 @@ class Queryable: yield row_pk, row @property - def columns(self) -> List["Column"]: + def columns(self) -> list["Column"]: "List of :ref:`Columns ` representing the columns in this table or view." if not self.exists(): return [] rows = self.db.execute( - "PRAGMA table_info({})".format(quote_identifier(self.name)) + f"PRAGMA table_info({quote_identifier(self.name)})" ).fetchall() return [Column(*row) for row in rows] @property - def columns_dict(self) -> Dict[str, Any]: + def columns_dict(self) -> dict[str, Any]: "``{column_name: python-type}`` dictionary representing columns in this table or view." return {column.name: column_affinity(column.type) for column in self.columns} @@ -2146,48 +2108,48 @@ class Table(Queryable): """ #: The ``rowid`` of the last inserted, updated or selected row. - last_rowid: Optional[int] = None + last_rowid: int | None = None #: The primary key of the last inserted, updated or selected row. - last_pk: Optional[Any] = None + last_pk: Any | None = None def __init__( self, db: Database, name: str, - pk: Optional[Any] = None, - foreign_keys: Optional[ForeignKeysType] = None, - column_order: Optional[List[str]] = None, - not_null: Optional[Iterable[str]] = None, - defaults: Optional[Dict[str, Any]] = None, + pk: Any | None = None, + foreign_keys: ForeignKeysType | None = None, + column_order: list[str] | None = None, + not_null: Iterable[str] | None = None, + defaults: dict[str, Any] | None = None, batch_size: int = 100, - hash_id: Optional[str] = None, - hash_id_columns: Optional[Iterable[str]] = None, + hash_id: str | None = None, + hash_id_columns: Iterable[str] | None = None, alter: bool = False, ignore: bool = False, replace: bool = False, - extracts: Optional[Union[Dict[str, str], List[str]]] = None, - conversions: Optional[dict] = None, - columns: Optional[Dict[str, Any]] = None, + extracts: dict[str, str] | list[str] | None = None, + conversions: dict | None = None, + columns: dict[str, Any] | None = None, strict: bool = False, ): super().__init__(db, name) - self._defaults = dict( - pk=pk, - foreign_keys=foreign_keys, - column_order=column_order, - not_null=not_null, - defaults=defaults, - batch_size=batch_size, - hash_id=hash_id, - hash_id_columns=hash_id_columns, - alter=alter, - ignore=ignore, - replace=replace, - extracts=extracts, - conversions=conversions or {}, - columns=columns, - strict=strict, - ) + self._defaults = { + "pk": pk, + "foreign_keys": foreign_keys, + "column_order": column_order, + "not_null": not_null, + "defaults": defaults, + "batch_size": batch_size, + "hash_id": hash_id, + "hash_id_columns": hash_id_columns, + "alter": alter, + "ignore": ignore, + "replace": replace, + "extracts": extracts, + "conversions": conversions or {}, + "columns": columns, + "strict": strict, + } def __repr__(self) -> str: return "".format( @@ -2212,7 +2174,7 @@ class Table(Queryable): return self.name in self.db.table_names() @property - def pks(self) -> List[str]: + def pks(self) -> list[str]: """ Primary key columns for this table, in PRIMARY KEY declaration order - ``PRAGMA table_info`` sets ``is_pk`` to the 1-based position of each @@ -2234,7 +2196,7 @@ class Table(Queryable): "Does this table use ``rowid`` for its primary key (no other primary keys are specified)?" return not any(column for column in self.columns if column.is_pk) - def get(self, pk_values: Union[list, tuple, str, int]) -> dict: + def get(self, pk_values: list | tuple | str | int) -> dict: """ Return row (as dictionary) for the specified primary key. @@ -2253,17 +2215,17 @@ class Table(Queryable): ) ) - wheres = ["{} = ?".format(quote_identifier(pk_name)) for pk_name in pks] + wheres = [f"{quote_identifier(pk_name)} = ?" for pk_name in pks] rows = self.rows_where(" and ".join(wheres), pk_values) try: - row = list(rows)[0] + row = next(iter(rows)) self.last_pk = last_pk return row - except IndexError: + except StopIteration: raise NotFoundError @property - def foreign_keys(self) -> List["ForeignKey"]: + def foreign_keys(self) -> list["ForeignKey"]: """ List of foreign keys defined on this table. @@ -2273,12 +2235,12 @@ class Table(Queryable): """ # PRAGMA foreign_key_list returns one row per column, grouped by "id" # with "seq" giving the column order within a compound foreign key. - by_id: Dict[int, list] = {} + by_id: dict[int, list] = {} for row in self.db.execute( - "PRAGMA foreign_key_list({})".format(quote_identifier(self.name)) + f"PRAGMA foreign_key_list({quote_identifier(self.name)})" ).fetchall(): if row is not None: - id, seq, table_name, from_, to_, on_update, on_delete, match = row + id, seq, table_name, from_, to_, on_update, on_delete, _match = row by_id.setdefault(id, []).append( (seq, table_name, from_, to_, on_update, on_delete) ) @@ -2311,7 +2273,7 @@ class Table(Queryable): return fks @property - def virtual_table_using(self) -> Optional[str]: + def virtual_table_using(self) -> str | None: "Type of virtual table, or ``None`` if this is not a virtual table." match = _virtual_table_using_re.match(self.schema) if match is None: @@ -2319,18 +2281,16 @@ class Table(Queryable): return match.groupdict()["using"].upper() @property - def indexes(self) -> List[Index]: + def indexes(self) -> list[Index]: "List of indexes defined on this table." - sql = 'PRAGMA index_list("{}")'.format(self.name) + sql = f'PRAGMA index_list("{self.name}")' indexes = [] for row in self.db.execute_returning_dicts(sql): index_name = row["name"] index_name_quoted = ( - '"{}"'.format(index_name) - if not index_name.startswith('"') - else index_name + f'"{index_name}"' if not index_name.startswith('"') else index_name ) - column_sql = "PRAGMA index_info({})".format(index_name_quoted) + column_sql = f"PRAGMA index_info({index_name_quoted})" columns = [] for seqno, cid, name in self.db.execute(column_sql).fetchall(): columns.append(name) @@ -2343,18 +2303,16 @@ class Table(Queryable): return indexes @property - def xindexes(self) -> List[XIndex]: + def xindexes(self) -> list[XIndex]: "List of indexes defined on this table using the more detailed ``XIndex`` format." - sql = 'PRAGMA index_list("{}")'.format(self.name) + sql = f'PRAGMA index_list("{self.name}")' indexes = [] for row in self.db.execute_returning_dicts(sql): index_name = row["name"] index_name_quoted = ( - '"{}"'.format(index_name) - if not index_name.startswith('"') - else index_name + f'"{index_name}"' if not index_name.startswith('"') else index_name ) - column_sql = "PRAGMA index_xinfo({})".format(index_name_quoted) + column_sql = f"PRAGMA index_xinfo({index_name_quoted})" index_columns = [] for info in self.db.execute(column_sql).fetchall(): index_columns.append(XIndexColumn(*info)) @@ -2362,7 +2320,7 @@ class Table(Queryable): return indexes @property - def triggers(self) -> List[Trigger]: + def triggers(self) -> list[Trigger]: "List of triggers defined on this table." return [ Trigger(*r) @@ -2374,12 +2332,12 @@ class Table(Queryable): ] @property - def triggers_dict(self) -> Dict[str, str]: + def triggers_dict(self) -> dict[str, str]: "``{trigger_name: sql}`` dictionary of triggers defined on this table." return {trigger.name: trigger.sql for trigger in self.triggers} @property - def default_values(self) -> Dict[str, Any]: + def default_values(self) -> dict[str, Any]: "``{column_name: default_value}`` dictionary of default values for columns in this table." return { column.name: _decode_default_value(column.default_value) @@ -2396,20 +2354,20 @@ class Table(Queryable): def create( self, - columns: Dict[str, Any], - pk: Optional[Any] = DEFAULT, - foreign_keys: Union[Optional[ForeignKeysType], Default] = DEFAULT, - column_order: Union[Optional[List[str]], Default] = DEFAULT, - not_null: Union[Optional[Iterable[str]], Default] = DEFAULT, - defaults: Union[Optional[Dict[str, Any]], Default] = DEFAULT, - hash_id: Union[Optional[str], Default] = DEFAULT, - hash_id_columns: Union[Optional[Iterable[str]], Default] = DEFAULT, - extracts: Union[Optional[Union[Dict[str, str], List[str]]], Default] = DEFAULT, + columns: dict[str, Any], + pk: Any | None = DEFAULT, + foreign_keys: ForeignKeysType | None | Default = DEFAULT, + column_order: list[str] | None | Default = DEFAULT, + not_null: Iterable[str] | None | Default = DEFAULT, + defaults: dict[str, Any] | None | Default = DEFAULT, + hash_id: str | None | Default = DEFAULT, + hash_id_columns: Iterable[str] | None | Default = DEFAULT, + extracts: dict[str, str] | list[str] | None | Default = DEFAULT, if_not_exists: bool = False, replace: bool = False, ignore: bool = False, transform: bool = False, - strict: Union[bool, Default] = DEFAULT, + strict: bool | Default = DEFAULT, ) -> "Table": """ Create a table with the specified columns. @@ -2493,28 +2451,25 @@ class Table(Queryable): if not self.exists(): raise NoTable(f"Table {self.name} does not exist") with self.db.atomic(): - sql = "CREATE TABLE {} AS SELECT * FROM {};".format( - quote_identifier(new_name), - quote_identifier(self.name), - ) + sql = f"CREATE TABLE {quote_identifier(new_name)} AS SELECT * FROM {quote_identifier(self.name)};" self.db.execute(sql) return self.db.table(new_name) def transform( self, *, - types: Optional[dict] = None, - rename: Optional[dict] = None, - drop: Optional[Iterable] = None, - pk: Optional[Any] = DEFAULT, - not_null: Optional[Iterable[str]] = None, - defaults: Optional[Dict[str, Any]] = None, - drop_foreign_keys: Optional[Iterable[str]] = None, - add_foreign_keys: Optional[ForeignKeysType] = None, - foreign_keys: Optional[ForeignKeysType] = None, - column_order: Optional[List[str]] = None, - keep_table: Optional[str] = None, - strict: Optional[bool] = None, + types: dict | None = None, + rename: dict | None = None, + drop: Iterable | None = None, + pk: Any | None = DEFAULT, + not_null: Iterable[str] | None = None, + defaults: dict[str, Any] | None = None, + drop_foreign_keys: Iterable[str] | None = None, + add_foreign_keys: ForeignKeysType | None = None, + foreign_keys: ForeignKeysType | None = None, + column_order: list[str] | None = None, + keep_table: str | None = None, + strict: bool | None = None, ) -> "Table": """ Apply an advanced alter table, including operations that are not supported by @@ -2633,20 +2588,20 @@ class Table(Queryable): def transform_sql( self, *, - types: Optional[dict] = None, - rename: Optional[dict] = None, - drop: Optional[Iterable] = None, - pk: Optional[Any] = DEFAULT, - not_null: Optional[Iterable[str]] = None, - defaults: Optional[Dict[str, Any]] = None, - drop_foreign_keys: Optional[Iterable] = None, - add_foreign_keys: Optional[ForeignKeysType] = None, - foreign_keys: Optional[ForeignKeysType] = None, - column_order: Optional[List[str]] = None, - tmp_suffix: Optional[str] = None, - keep_table: Optional[str] = None, - strict: Optional[bool] = None, - ) -> List[str]: + types: dict | None = None, + rename: dict | None = None, + drop: Iterable | None = None, + pk: Any | None = DEFAULT, + not_null: Iterable[str] | None = None, + defaults: dict[str, Any] | None = None, + drop_foreign_keys: Iterable | None = None, + add_foreign_keys: ForeignKeysType | None = None, + foreign_keys: ForeignKeysType | None = None, + column_order: list[str] | None = None, + tmp_suffix: str | None = None, + keep_table: str | None = None, + strict: bool | None = None, + ) -> list[str]: """ Return a list of SQL statements that should be executed in order to apply this transformation. @@ -2689,7 +2644,7 @@ class Table(Queryable): if isinstance(not_null, dict): not_null = { resolve_casing(c, existing_columns): v - for c, v in cast(Dict[str, Any], not_null).items() + for c, v in cast(dict[str, Any], not_null).items() } elif isinstance(not_null, set): not_null = {resolve_casing(c, existing_columns) for c in not_null} @@ -2700,7 +2655,7 @@ class Table(Queryable): if column_order is not None: column_order = [resolve_casing(c, existing_columns) for c in column_order] - create_table_foreign_keys: List[ForeignKeyIndicator] = [] + create_table_foreign_keys: list[ForeignKeyIndicator] = [] if foreign_keys is not None: if add_foreign_keys is not None: @@ -2777,9 +2732,7 @@ class Table(Queryable): for fk in self.db.resolve_foreign_keys(self.name, add_foreign_keys): create_table_foreign_keys.append(fk_with_renamed_columns(fk)) - new_table_name = "{}_new_{}".format( - self.name, tmp_suffix or os.urandom(6).hex() - ) + new_table_name = f"{self.name}_new_{tmp_suffix or os.urandom(6).hex()}" current_column_pairs = list(self.columns_dict.items()) new_column_pairs = [] copy_from_to = {column: column for column, _ in current_column_pairs} @@ -2824,9 +2777,7 @@ class Table(Queryable): pass else: raise ValueError( - "not_null must be a dict or a set or None, it was {}".format( - repr(not_null) - ) + f"not_null must be a dict or a set or None, it was {not_null!r}" ) # defaults= create_table_defaults = { @@ -2876,17 +2827,13 @@ class Table(Queryable): # Drop (or keep) the old table if keep_table: sqls.append( - "ALTER TABLE {} RENAME TO {};".format( - quote_identifier(self.name), quote_identifier(keep_table) - ) + f"ALTER TABLE {quote_identifier(self.name)} RENAME TO {quote_identifier(keep_table)};" ) else: - sqls.append("DROP TABLE {};".format(quote_identifier(self.name))) + sqls.append(f"DROP TABLE {quote_identifier(self.name)};") # Rename the new one sqls.append( - "ALTER TABLE {} RENAME TO {};".format( - quote_identifier(new_table_name), quote_identifier(self.name) - ) + f"ALTER TABLE {quote_identifier(new_table_name)} RENAME TO {quote_identifier(self.name)};" ) # Re-add existing indexes for index in self.indexes: @@ -2904,7 +2851,7 @@ class Table(Queryable): if keep_table: sqls.append(f"DROP INDEX IF EXISTS {quote_identifier(index.name)};") for col in index.columns: - if col in rename.keys() or col in drop: + if col in rename or col in drop: raise TransformError( f"Index '{index.name}' column '{col}' is not in updated table '{self.name}'. " f"You must manually drop this index prior to running this transformation " @@ -2916,10 +2863,10 @@ class Table(Queryable): def extract( self, - columns: Union[str, Iterable[str]], - table: Optional[str] = None, - fk_column: Optional[str] = None, - rename: Optional[Dict[str, str]] = None, + columns: str | Iterable[str], + table: str | None = None, + fk_column: str | None = None, + rename: dict[str, str] | None = None, ) -> "Table": """ Extract specified columns into a separate table. @@ -2938,15 +2885,13 @@ class Table(Queryable): rename = {resolve_casing(k, self.columns_dict): v for k, v in rename.items()} if not set(columns).issubset(self.columns_dict.keys()): raise InvalidColumns( - "Invalid columns {} for table with columns {}".format( - columns, list(self.columns_dict.keys()) - ) + f"Invalid columns {columns} for table with columns {list(self.columns_dict.keys())}" ) with self.db.atomic(): table = table or "_".join(columns) lookup_table = self.db.table(table) - fk_column = fk_column or "{}_id".format(table) - magic_lookup_column = "{}_{}".format(fk_column, os.urandom(6).hex()) + fk_column = fk_column or f"{table}_id" + magic_lookup_column = f"{fk_column}_{os.urandom(6).hex()}" # Populate the lookup table with all of the extracted unique values lookup_columns_definition = { @@ -2959,16 +2904,12 @@ class Table(Queryable): lookup_table.columns_dict.items() ): raise InvalidColumns( - "Lookup table {} already exists but does not have columns {}".format( - table, lookup_columns_definition - ) + f"Lookup table {table} already exists but does not have columns {lookup_columns_definition}" ) else: lookup_table.create( { - **{ - "id": int, - }, + "id": int, **lookup_columns_definition, }, pk="id", @@ -2978,19 +2919,14 @@ class Table(Queryable): # Rows where every extracted column is null are left alone - they # get a null foreign key and no lookup table record, see #186 all_columns_are_null = " AND ".join( - "{} IS NULL".format(quote_identifier(c)) for c in columns + f"{quote_identifier(c)} IS NULL" for c in columns ) # INSERT OR IGNORE dedupes against the unique index, but unique # indexes treat NULLs as distinct - the NOT EXISTS guard uses IS # comparison so NULL-containing rows match existing lookup rows # instead of being inserted again already_in_lookup = " AND ".join( - "{lookup}.{lookup_col} IS {source}.{source_col}".format( - lookup=quote_identifier(table), - lookup_col=quote_identifier(rename.get(column) or column), - source=quote_identifier(self.name), - source_col=quote_identifier(column), - ) + f"{quote_identifier(table)}.{quote_identifier(rename.get(column) or column)} IS {quote_identifier(self.name)}.{quote_identifier(column)}" for column in columns ) self.db.execute( @@ -3018,12 +2954,10 @@ class Table(Queryable): quote_identifier(magic_lookup_column), quote_identifier(table), where=" AND ".join( - "{}.{} IS {}.{}".format( - quote_identifier(self.name), - quote_identifier(column), - quote_identifier(table), - quote_identifier(rename.get(column) or column), - ) + f"{quote_identifier(self.name)}." + f"{quote_identifier(column)} IS " + f"{quote_identifier(table)}." + f"{quote_identifier(rename.get(column) or column)}" for column in columns ), all_null=all_columns_are_null, @@ -3052,8 +2986,8 @@ class Table(Queryable): def create_index( self, - columns: Iterable[Union[str, DescIndex]], - index_name: Optional[str] = None, + columns: Iterable[str | DescIndex], + index_name: str | None = None, unique: bool = False, if_not_exists: bool = False, find_unique_name: bool = False, @@ -3080,16 +3014,14 @@ class Table(Queryable): columns_sql = [] for column in columns: if isinstance(column, DescIndex): - columns_sql.append("{} desc".format(quote_identifier(column))) + columns_sql.append(f"{quote_identifier(column)} desc") else: columns_sql.append(quote_identifier(column)) suffix = None created_index_name = None while True: - created_index_name = ( - "{}_{}".format(index_name, suffix) if suffix else index_name - ) + created_index_name = f"{index_name}_{suffix}" if suffix else index_name sql = ( textwrap.dedent(""" CREATE {unique}INDEX {if_not_exists}{index_name} @@ -3121,7 +3053,7 @@ class Table(Queryable): suffix += 1 continue else: - raise e + raise if analyze: self.db.analyze(created_index_name) return self @@ -3136,19 +3068,17 @@ class Table(Queryable): if index_name not in {index.name for index in self.indexes}: if ignore: return self - raise OperationalError( - "No index named {} on table {}".format(index_name, self.name) - ) - self.db.execute("DROP INDEX {}".format(quote_identifier(index_name))) + raise OperationalError(f"No index named {index_name} on table {self.name}") + self.db.execute(f"DROP INDEX {quote_identifier(index_name)}") return self def add_column( self, col_name: str, - col_type: Optional[Any] = None, - fk: Optional[str] = None, - fk_col: Optional[str] = None, - not_null_default: Optional[Any] = None, + col_type: Any | None = None, + fk: str | None = None, + fk_col: str | None = None, + not_null_default: Any | None = None, ): """ Add a column to this table. See :ref:`python_api_add_column`. @@ -3163,12 +3093,12 @@ class Table(Queryable): if fk is not None: # fk must be a valid table if fk not in self.db.table_names(): - raise AlterError("table '{}' does not exist".format(fk)) + raise AlterError(f"table '{fk}' does not exist") # if fk_col specified, must be a valid column if fk_col is not None: fk_col = resolve_casing(fk_col, self.db[fk].columns_dict) if fk_col not in self.db[fk].columns_dict: - raise AlterError("table '{}' has no column {}".format(fk, fk_col)) + raise AlterError(f"table '{fk}' has no column {fk_col}") else: # automatically set fk_col to first primary_key of fk table pks = sorted( @@ -3185,8 +3115,8 @@ class Table(Queryable): col_type = str not_null_sql = None if not_null_default is not None: - not_null_sql = "NOT NULL DEFAULT {}".format( - self.db.quote_default_value(not_null_default) + not_null_sql = ( + f"NOT NULL DEFAULT {self.db.quote_default_value(not_null_default)}" ) sql = "ALTER TABLE {} ADD COLUMN {} {col_type}{not_null_default};".format( quote_identifier(self.name), @@ -3206,7 +3136,7 @@ class Table(Queryable): :param ignore: Set to ``True`` to ignore the error if the table does not exist """ try: - self.db.execute("DROP TABLE {}".format(quote_identifier(self.name))) + self.db.execute(f"DROP TABLE {quote_identifier(self.name)}") except sqlite3.OperationalError: if not ignore: raise @@ -3238,16 +3168,14 @@ class Table(Queryable): return existing_tables[table] # If we get here there's no obvious candidate - raise an error raise NoObviousTable( - "No obvious foreign key table for column '{}' - tried {}".format( - column, repr(possibilities) - ) + f"No obvious foreign key table for column '{column}' - tried {possibilities!r}" ) def guess_foreign_column(self, other_table: str) -> str: pks = [c for c in self.db[other_table].columns if c.is_pk] if len(pks) != 1: raise BadPrimaryKey( - "Could not detect single primary key for table '{}'".format(other_table) + f"Could not detect single primary key for table '{other_table}'" ) else: return pks[0].name @@ -3255,8 +3183,8 @@ class Table(Queryable): def add_foreign_key( self, column: ForeignKeyColumns, - other_table: Optional[str] = None, - other_column: Optional[ForeignKeyColumns] = None, + other_table: str | None = None, + other_column: ForeignKeyColumns | None = None, ignore: bool = False, on_delete: str = "NO ACTION", on_update: str = "NO ACTION", @@ -3279,7 +3207,7 @@ class Table(Queryable): # Ensure columns exist for col in columns: if col not in self.columns_dict: - raise AlterError("No such column: {}".format(col)) + raise AlterError(f"No such column: {col}") # If other_table is not specified, attempt to guess it from the column if other_table is None: if len(columns) > 1: @@ -3312,7 +3240,7 @@ class Table(Queryable): not [c for c in self.db[other_table].columns if c.name == other_col] and other_col != "rowid" ): - raise AlterError("No such column: {}.{}".format(other_table, other_col)) + raise AlterError(f"No such column: {other_table}.{other_col}") # Check we do not already have an existing foreign key if any( fk @@ -3413,9 +3341,7 @@ class Table(Queryable): def has_counts_triggers(self) -> bool: "Does this table have triggers setup to update cached counts?" trigger_names = { - "{table}{counts_table}_{suffix}".format( - counts_table=self.db._counts_table_name, table=self.name, suffix=suffix - ) + f"{self.name}{self.db._counts_table_name}_{suffix}" for suffix in ["insert", "delete"] } return trigger_names.issubset(self.triggers_dict.keys()) @@ -3425,7 +3351,7 @@ class Table(Queryable): columns: Iterable[str], fts_version: str = "FTS5", create_triggers: bool = False, - tokenize: Optional[str] = None, + tokenize: str | None = None, replace: bool = False, ): """ @@ -3452,13 +3378,13 @@ class Table(Queryable): table_fts=quote_identifier(self.name + "_fts"), columns=", ".join(quote_identifier(c) for c in columns), fts_version=fts_version, - tokenize="\n tokenize='{}',".format(tokenize) if tokenize else "", + tokenize=f"\n tokenize='{tokenize}'," if tokenize else "", ) ) should_recreate = False - if replace and self.db["{}_fts".format(self.name)].exists(): + if replace and self.db[f"{self.name}_fts"].exists(): # Does the table need to be recreated? - fts_schema = self.db["{}_fts".format(self.name)].schema + fts_schema = self.db[f"{self.name}_fts"].schema if fts_schema != create_fts_sql: should_recreate = True expected_triggers = {self.name + suffix for suffix in ("_ai", "_ad", "_au")} @@ -3477,8 +3403,8 @@ class Table(Queryable): self.populate_fts(columns) if create_triggers: - old_cols = ", ".join("old.{}".format(quote_identifier(c)) for c in columns) - new_cols = ", ".join("new.{}".format(quote_identifier(c)) for c in columns) + old_cols = ", ".join(f"old.{quote_identifier(c)}" for c in columns) + new_cols = ", ".join(f"new.{quote_identifier(c)}" for c in columns) columns_quoted = ", ".join(quote_identifier(c) for c in columns) table = quote_identifier(self.name) table_fts = quote_identifier(self.name + "_fts") @@ -3550,7 +3476,7 @@ class Table(Queryable): with self.db.atomic(): for trigger_name in trigger_names: self.db.execute( - "DROP TRIGGER IF EXISTS {}".format(quote_identifier(trigger_name)) + f"DROP TRIGGER IF EXISTS {quote_identifier(trigger_name)}" ) return self @@ -3568,7 +3494,7 @@ class Table(Queryable): ) return self - def detect_fts(self) -> Optional[str]: + def detect_fts(self) -> str | None: "Detect if table has a corresponding FTS virtual table and return it" sql = textwrap.dedent(""" SELECT name FROM sqlite_master @@ -3583,8 +3509,8 @@ class Table(Queryable): ) """).strip() args = { - "like": "%VIRTUAL TABLE%USING FTS%content=[{}]%".format(self.name), - "like2": '%VIRTUAL TABLE%USING FTS%content="{}"%'.format(self.name), + "like": f"%VIRTUAL TABLE%USING FTS%content=[{self.name}]%", + "like2": f'%VIRTUAL TABLE%USING FTS%content="{self.name}"%', "table": self.name, } rows = self.db.execute(sql, args).fetchall() @@ -3605,11 +3531,11 @@ class Table(Queryable): def search_sql( self, - columns: Optional[Iterable[str]] = None, - order_by: Optional[str] = None, - limit: Optional[int] = None, - offset: Optional[int] = None, - where: Optional[str] = None, + columns: Iterable[str] | None = None, + order_by: str | None = None, + limit: int | None = None, + offset: int | None = None, + where: str | None = None, include_rank: bool = False, ) -> str: """ " @@ -3626,16 +3552,16 @@ class Table(Queryable): original = "original_" if self.name == "original" else "original" original_quoted = quote_identifier(original) columns_sql = "*" - columns_with_prefix_sql = "{}.*".format(original_quoted) + columns_with_prefix_sql = f"{original_quoted}.*" if columns: columns_sql = ",\n ".join(quote_identifier(c) for c in columns) columns_with_prefix_sql = ",\n ".join( - "{}.{}".format(original_quoted, quote_identifier(c)) for c in columns + f"{original_quoted}.{quote_identifier(c)}" for c in columns ) fts_table = self.detect_fts() if not fts_table: raise ValueError( - "Full-text search is not configured for table '{}'".format(self.name) + f"Full-text search is not configured for table '{self.name}'" ) fts_table_quoted = quote_identifier(fts_table) virtual_table_using = self.db.table(fts_table).virtual_table_using @@ -3658,22 +3584,20 @@ class Table(Queryable): {limit_offset} """).strip() if virtual_table_using == "FTS5": - rank_implementation = "{}.rank".format(fts_table_quoted) + rank_implementation = f"{fts_table_quoted}.rank" else: self.db.register_fts4_bm25() - rank_implementation = "rank_bm25(matchinfo({}, 'pcnalx'))".format( - fts_table_quoted - ) + rank_implementation = f"rank_bm25(matchinfo({fts_table_quoted}, 'pcnalx'))" if include_rank: columns_with_prefix_sql += ",\n " + rank_implementation + " rank" limit_offset = "" if limit is not None: - limit_offset += " limit {}".format(limit) + limit_offset += f" limit {limit}" if offset is not None: - limit_offset += " offset {}".format(offset) + limit_offset += f" offset {offset}" return sql.format( dbtable=quote_identifier(self.name), - where_clause="\n where {}".format(where) if where else "", + where_clause=f"\n where {where}" if where else "", original=original_quoted, columns=columns_sql, columns_with_prefix=columns_with_prefix_sql, @@ -3685,12 +3609,12 @@ class Table(Queryable): def search( self, q: str, - order_by: Optional[str] = None, - columns: Optional[Iterable[str]] = None, - limit: Optional[int] = None, - offset: Optional[int] = None, - where: Optional[str] = None, - where_args: Optional[Union[Iterable, dict]] = None, + order_by: str | None = None, + columns: Iterable[str] | None = None, + limit: int | None = None, + offset: int | None = None, + where: str | None = None, + where_args: Iterable | dict | None = None, include_rank: bool = False, quote: bool = False, ) -> Generator[dict, None, None]: @@ -3736,7 +3660,7 @@ class Table(Queryable): def value_or_default(self, key: str, value: Any) -> Any: 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: list | tuple | str | float) -> "Table": """ Delete row matching the specified primary key. @@ -3745,7 +3669,7 @@ class Table(Queryable): if not isinstance(pk_values, (list, tuple)): pk_values = [pk_values] self.get(pk_values) - wheres = ["{} = ?".format(quote_identifier(pk_name)) for pk_name in self.pks] + wheres = [f"{quote_identifier(pk_name)} = ?" for pk_name in self.pks] sql = "delete from {} where {wheres}".format( quote_identifier(self.name), wheres=" and ".join(wheres) ) @@ -3755,8 +3679,8 @@ class Table(Queryable): def delete_where( self, - where: Optional[str] = None, - where_args: Optional[Union[Sequence, Dict[str, Any]]] = None, + where: str | None = None, + where_args: Sequence | dict[str, Any] | None = None, analyze: bool = False, ) -> "Table": """ @@ -3771,7 +3695,7 @@ class Table(Queryable): """ if not self.exists(): return self - sql = "delete from {}".format(quote_identifier(self.name)) + sql = f"delete from {quote_identifier(self.name)}" if where is not None: sql += " where " + where with self.db.atomic(): @@ -3782,10 +3706,10 @@ class Table(Queryable): def update( self, - pk_values: Union[list, tuple, str, int, float], - updates: Optional[dict] = None, + pk_values: list | tuple | str | float, + updates: dict | None = None, alter: bool = False, - conversions: Optional[dict] = None, + conversions: dict | None = None, ) -> "Table": """ Execute a SQL ``UPDATE`` against the specified row. @@ -3816,7 +3740,7 @@ class Table(Queryable): "{} = {}".format(quote_identifier(key), conversions.get(key, "?")) ) args.append(jsonify_if_needed(value)) - wheres = ["{} = ?".format(quote_identifier(pk_name)) for pk_name in pks] + wheres = [f"{quote_identifier(pk_name)} = ?" for pk_name in pks] args.extend(pk_values) sql = "update {} set {sets} where {wheres}".format( quote_identifier(self.name), @@ -3841,14 +3765,14 @@ class Table(Queryable): def convert( self, - columns: Union[str, List[str]], + columns: str | list[str], fn: Callable, - output: Optional[str] = None, - output_type: Optional[Any] = None, + output: str | None = None, + output_type: Any | None = None, drop: bool = False, multi: bool = False, - where: Optional[str] = None, - where_args: Optional[Union[Sequence, Dict[str, Any]]] = None, + where: str | None = None, + where_args: Sequence | dict[str, Any] | None = None, show_progress: bool = False, ) -> "Table": """ @@ -3905,15 +3829,11 @@ class Table(Queryable): quote_identifier(self.name), sets=", ".join( [ - "{} = {}({})".format( - quote_identifier(output or column), - fn_name, - quote_identifier(column), - ) + f"{quote_identifier(output or column)} = {fn_name}({quote_identifier(column)})" for column in columns ] ), - where=" where {}".format(where) if where is not None else "", + where=f" where {where}" if where is not None else "", ) with self.db.atomic(): self.db.execute(sql, where_args or []) @@ -3926,7 +3846,7 @@ class Table(Queryable): ): # First we execute the function pk_to_values = {} - new_column_types: Dict[str, Set[type]] = {} + new_column_types: dict[str, set[type]] = {} pks = self.pks with progressbar( @@ -3958,15 +3878,17 @@ class Table(Queryable): self.add_column(column_name, column_type) # Run the updates - with progressbar( - length=self.count, silent=not show_progress, label="2: Updating" - ) as bar: - with self.db.atomic(): - for pk, updates in pk_to_values.items(): - self.update(pk, updates) - bar.update(1) - if drop: - self.transform(drop=(column,)) + with ( + progressbar( + length=self.count, silent=not show_progress, label="2: Updating" + ) as bar, + self.db.atomic(), + ): + for pk, updates in pk_to_values.items(): + self.update(pk, updates) + bar.update(1) + if drop: + self.transform(drop=(column,)) def build_insert_queries_and_params( self, @@ -4166,9 +4088,7 @@ class Table(Queryable): ) for col in set_cols ), - wheres=" AND ".join( - "{} = ?".format(quote_identifier(pk)) for pk in pks - ), + wheres=" AND ".join(f"{quote_identifier(pk)} = ?" for pk in pks), ) queries_and_params.append( ( @@ -4201,7 +4121,7 @@ class Table(Queryable): replace, ignore, list_mode=False, - ) -> Optional[sqlite3.Cursor]: + ) -> sqlite3.Cursor | None: queries_and_params = self.build_insert_queries_and_params( extracts, chunk, @@ -4271,21 +4191,21 @@ class Table(Queryable): def insert( self, - record: Dict[str, Any], + record: dict[str, Any], pk=DEFAULT, foreign_keys=DEFAULT, - column_order: Optional[Union[List[str], Default]] = DEFAULT, - not_null: Optional[Union[Iterable[str], Default]] = DEFAULT, - defaults: Optional[Union[Dict[str, Any], Default]] = DEFAULT, - hash_id: Optional[Union[str, Default]] = DEFAULT, - hash_id_columns: Optional[Union[Iterable[str], Default]] = DEFAULT, - alter: Optional[Union[bool, Default]] = DEFAULT, - ignore: Optional[Union[bool, Default]] = DEFAULT, - replace: Optional[Union[bool, Default]] = DEFAULT, - extracts: Optional[Union[Dict[str, str], List[str], Default]] = DEFAULT, - conversions: Optional[Union[Dict[str, str], Default]] = DEFAULT, - columns: Optional[Union[Dict[str, Any], Default]] = DEFAULT, - strict: Optional[Union[bool, Default]] = DEFAULT, + column_order: list[str] | Default | None = DEFAULT, + not_null: Iterable[str] | Default | None = DEFAULT, + defaults: dict[str, Any] | Default | None = DEFAULT, + hash_id: str | Default | None = DEFAULT, + hash_id_columns: Iterable[str] | Default | None = DEFAULT, + alter: bool | Default | None = DEFAULT, + ignore: bool | Default | None = DEFAULT, + replace: bool | Default | None = DEFAULT, + extracts: dict[str, str] | list[str] | Default | None = DEFAULT, + conversions: dict[str, str] | Default | None = DEFAULT, + columns: dict[str, Any] | Default | None = DEFAULT, + strict: bool | Default | None = DEFAULT, ) -> "Table": """ Insert a single record into the table. The table will be created with a schema that matches @@ -4340,10 +4260,7 @@ class Table(Queryable): def insert_all( self, - records: Union[ - Iterable[Dict[str, Any]], - Iterable[Sequence[Any]], - ], + records: Iterable[dict[str, Any]] | Iterable[Sequence[Any]], pk=DEFAULT, foreign_keys=DEFAULT, column_order=DEFAULT, @@ -4440,7 +4357,7 @@ class Table(Queryable): # Detect if we're using list-based iteration or dict-based iteration list_mode = False - column_names: List[str] = [] + column_names: list[str] = [] # Fix up any records with square braces in the column names (only for dict mode) # We'll handle this differently for list mode @@ -4460,7 +4377,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 = cast(List[str], 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 @@ -4469,7 +4386,7 @@ class Table(Queryable): except StopIteration: return self # Only headers, no data if not isinstance(first_record, (list, tuple)): - raise ValueError( + raise ValueError( # noqa: TRY004 "After column names list, all subsequent records must also be lists" ) else: @@ -4479,13 +4396,11 @@ class Table(Queryable): first_record = next(records_iter) except StopIteration: return self - first_record = cast(Dict[str, Any], first_record) + first_record = cast(dict[str, Any], first_record) num_columns = len(first_record.keys()) if num_columns > SQLITE_MAX_VARS: - raise ValueError( - "Rows can have a maximum of {} columns".format(SQLITE_MAX_VARS) - ) + raise ValueError(f"Rows can have a maximum of {SQLITE_MAX_VARS} columns") batch_size = ( 1 if num_columns == 0 @@ -4495,7 +4410,7 @@ class Table(Queryable): self.last_pk = None if truncate and self.exists(): with self.db.atomic(): - self.db.execute("DELETE FROM {};".format(quote_identifier(self.name))) + self.db.execute(f"DELETE FROM {quote_identifier(self.name)};") result = None for chunk in chunks(itertools.chain([first_record], records_iter), batch_size): chunk = list(chunk) @@ -4508,7 +4423,7 @@ class Table(Queryable): chunk_as_dicts = [dict(zip(column_names, row)) for row in chunk] column_types = suggest_column_types(chunk_as_dicts) else: - dict_chunk = cast(List[Dict[str, Any]], chunk) + dict_chunk = cast(list[dict[str, Any]], chunk) column_types = suggest_column_types(dict_chunk) if extracts: for col in extracts: @@ -4535,10 +4450,10 @@ class Table(Queryable): if hash_id: all_columns.insert(0, hash_id) else: - all_columns_set: Set[str] = set() - for record in cast(List[Dict[str, Any]], chunk): + all_columns_set: set[str] = set() + for record in cast(list[dict[str, Any]], chunk): all_columns_set.update(record.keys()) - all_columns = list(sorted(all_columns_set)) + all_columns = sorted(all_columns_set) if hash_id: all_columns.insert(0, hash_id) if deferred_invalid_pk_check is not None: @@ -4553,7 +4468,7 @@ class Table(Queryable): raise invalid_pk_error else: if not list_mode: - for record in cast(List[Dict[str, Any]], chunk): + for record in cast(list[dict[str, Any]], chunk): all_columns += [ column for column in record if column not in all_columns ] @@ -4592,7 +4507,7 @@ class Table(Queryable): zip(column_names, cast(Sequence[Any], first_record)) ) else: - first_record_dict = cast(Dict[str, Any], first_record) + first_record_dict = cast(dict[str, Any], first_record) if hash_id: self.last_pk = hash_record(first_record_dict, hash_id_columns) elif isinstance(pk, str): @@ -4608,7 +4523,7 @@ class Table(Queryable): # columns so we can report its rowid (and pk if not already # known). Falls back to leaving them unset if the conflict # cannot be resolved to a pk lookup (e.g. a UNIQUE column). - key_cols: Optional[List[str]] = None + key_cols: list[str] | None = None if isinstance(pk, str): key_cols = [pk] elif pk: @@ -4625,12 +4540,10 @@ class Table(Queryable): key_values = None if key_values is not None: where = " and ".join( - "{} = ?".format(quote_identifier(c)) for c in key_cols + f"{quote_identifier(c)} = ?" for c in key_cols ) existing = self.db.execute( - "select rowid from {} where {} limit 1".format( - quote_identifier(self.name), where - ), + f"select rowid from {quote_identifier(self.name)} where {where} limit 1", key_values, ).fetchone() if existing is not None: @@ -4650,7 +4563,9 @@ class Table(Queryable): rowid_pk = isinstance(pk, str) and pk.lower() in ROWID_ALIASES if (hash_id or (pk and not rowid_pk)) and self.last_rowid: # Set self.last_pk to the pk(s) for that rowid - row = list(self.rows_where("rowid = ?", [self.last_rowid]))[0] + row = next( + iter(self.rows_where("rowid = ?", [self.last_rowid])) + ) if hash_id: self.last_pk = row[hash_id] elif isinstance(pk, str): @@ -4680,7 +4595,7 @@ class Table(Queryable): for p in pk ) else: - first_record_dict = cast(Dict[str, Any], first_record) + first_record_dict = cast(dict[str, Any], first_record) if hash_id: self.last_pk = hash_record(first_record_dict, hash_id_columns) else: @@ -4738,10 +4653,7 @@ class Table(Queryable): def upsert_all( self, - records: Union[ - Iterable[Dict[str, Any]], - Iterable[Sequence[Any]], - ], + records: Iterable[dict[str, Any]] | Iterable[Sequence[Any]], pk=DEFAULT, foreign_keys=DEFAULT, column_order=DEFAULT, @@ -4779,7 +4691,7 @@ class Table(Queryable): strict=strict, ) - def add_missing_columns(self, records: Iterable[Dict[str, Any]]) -> "Table": + def add_missing_columns(self, records: Iterable[dict[str, Any]]) -> "Table": needed_columns = suggest_column_types(records) current_columns = {c.lower() for c in self.columns_dict} for col_name, col_type in needed_columns.items(): @@ -4789,17 +4701,17 @@ class Table(Queryable): def lookup( self, - lookup_values: Dict[str, Any], - extra_values: Optional[Dict[str, Any]] = None, - pk: Optional[str] = "id", - foreign_keys: Optional[ForeignKeysType] = None, - column_order: Optional[List[str]] = None, - not_null: Optional[Iterable[str]] = None, - defaults: Optional[Dict[str, Any]] = None, - extracts: Optional[Union[Dict[str, str], List[str]]] = None, - conversions: Optional[Dict[str, str]] = None, - columns: Optional[Dict[str, Any]] = None, - strict: Optional[bool] = False, + lookup_values: dict[str, Any], + extra_values: dict[str, Any] | None = None, + pk: str | None = "id", + foreign_keys: ForeignKeysType | None = None, + column_order: list[str] | None = None, + not_null: Iterable[str] | None = None, + defaults: dict[str, Any] | None = None, + extracts: dict[str, str] | list[str] | None = None, + conversions: dict[str, str] | None = None, + columns: dict[str, Any] | None = None, + strict: bool | None = False, ): """ Create or populate a lookup table with the specified values. @@ -4825,7 +4737,7 @@ class Table(Queryable): :param strict: Boolean, apply STRICT mode if creating the table. """ if not isinstance(lookup_values, dict): - raise ValueError("lookup_values must be a dictionary") + raise ValueError("lookup_values must be a dictionary") # noqa: TRY004 if pk is None: raise ValueError("pk cannot be None") if extra_values is not None and not isinstance(extra_values, dict): @@ -4843,9 +4755,7 @@ class Table(Queryable): } not in unique_column_sets: self.create_index(lookup_values.keys(), unique=True) # IS rather than = so that null values are matched correctly - wheres = [ - "{} IS ?".format(quote_identifier(column)) for column in lookup_values - ] + wheres = [f"{quote_identifier(column)} IS ?" for column in lookup_values] rows = list( self.rows_where( " and ".join(wheres), [value for _, value in lookup_values.items()] @@ -4885,12 +4795,10 @@ class Table(Queryable): def m2m( self, other_table: Union[str, "Table"], - record_or_iterable: Optional[ - Union[Iterable[Dict[str, Any]], Dict[str, Any]] - ] = None, - pk: Optional[Union[Any, Default]] = DEFAULT, - lookup: Optional[Dict[str, Any]] = None, - m2m_table: Optional[str] = None, + record_or_iterable: Iterable[dict[str, Any]] | dict[str, Any] | None = None, + pk: Any | Default | None = DEFAULT, + lookup: dict[str, Any] | None = None, + m2m_table: str | None = None, alter: bool = False, ): """ @@ -4923,8 +4831,8 @@ class Table(Queryable): raise ValueError("Provide lookup= or record, not both") elif record_or_iterable is None: raise ValueError("Provide lookup= or record, not both") - tables = list(sorted([self.name, other_table.name])) - columns = ["{}_id".format(t) for t in tables] + tables = sorted([self.name, other_table.name]) + columns = [f"{t}_id" for t in tables] if m2m_table is not None: m2m_table_name = m2m_table else: @@ -4934,9 +4842,7 @@ class Table(Queryable): m2m_table_name = candidates[0] elif len(candidates) > 1: raise NoObviousTable( - "No single obvious m2m table for {}, {} - use m2m_table= parameter".format( - self.name, other_table.name - ) + f"No single obvious m2m table for {self.name}, {other_table.name} - use m2m_table= parameter" ) else: # If not, create a new table @@ -4947,7 +4853,7 @@ class Table(Queryable): if isinstance(record_or_iterable, Mapping): records = [record_or_iterable] else: - records = cast(List, record_or_iterable) + records = cast(list, record_or_iterable) # Ensure each record exists in other table for record in records: id = other_table.insert( @@ -4955,8 +4861,8 @@ class Table(Queryable): ).last_pk m2m_table_obj.insert( { - "{}_id".format(other_table.name): id, - "{}_id".format(self.name): our_id, + f"{other_table.name}_id": id, + f"{self.name}_id": our_id, }, replace=True, ) @@ -4964,8 +4870,8 @@ class Table(Queryable): id = other_table.lookup(lookup) m2m_table_obj.insert( { - "{}_id".format(other_table.name): id, - "{}_id".format(self.name): our_id, + f"{other_table.name}_id": id, + f"{self.name}_id": our_id, }, replace=True, ) @@ -5012,21 +4918,19 @@ class Table(Queryable): table_quoted = quote_identifier(table) column_quoted = quote_identifier(column) num_null = db.execute( - "select count(*) from {} where {} is null".format( - table_quoted, column_quoted - ) + f"select count(*) from {table_quoted} where {column_quoted} is null" ).fetchone()[0] num_blank = db.execute( - "select count(*) from {} where {} = ''".format(table_quoted, column_quoted) + f"select count(*) from {table_quoted} where {column_quoted} = ''" ).fetchone()[0] num_distinct = db.execute( - "select count(distinct {}) from {}".format(column_quoted, table_quoted) + f"select count(distinct {column_quoted}) from {table_quoted}" ).fetchone()[0] most_common_results = None least_common_results = None if num_distinct == 1: value = db.execute( - "select {} from {} limit 1".format(column_quoted, table_quoted) + f"select {column_quoted} from {table_quoted} limit 1" ).fetchone()[0] most_common_results = [(truncate(value), total_rows)] elif num_distinct != total_rows: @@ -5038,13 +4942,10 @@ class Table(Queryable): most_common_results = [ (truncate(r[0]), r[1]) for r in db.execute( - "select {}, count(*) from {} group by {} order by count(*) desc, {} limit {}".format( - column_quoted, - table_quoted, - column_quoted, - column_quoted, - common_limit, - ) + f"select {column_quoted}, count(*) " + f"from {table_quoted} group by {column_quoted} " + f"order by count(*) desc, {column_quoted} " + f"limit {common_limit}" ).fetchall() ] most_common_results.sort(key=lambda p: (p[1], p[0]), reverse=True) @@ -5056,13 +4957,10 @@ class Table(Queryable): least_common_results = [ (truncate(r[0]), r[1]) for r in db.execute( - "select {}, count(*) from {} group by {} order by count(*), {} desc limit {}".format( - column_quoted, - table_quoted, - column_quoted, - column_quoted, - common_limit, - ) + f"select {column_quoted}, count(*) " + f"from {table_quoted} group by {column_quoted} " + f"order by count(*), {column_quoted} desc " + f"limit {common_limit}" ).fetchall() ] least_common_results.sort(key=lambda p: (p[1], p[0])) @@ -5179,7 +5077,7 @@ class View(Queryable): """ try: - self.db.execute("DROP VIEW {}".format(quote_identifier(self.name))) + self.db.execute(f"DROP VIEW {quote_identifier(self.name)}") except sqlite3.OperationalError: if not ignore: raise @@ -5192,16 +5090,14 @@ def jsonify_if_needed(value: object) -> object: return json.dumps(value, default=repr, ensure_ascii=False) elif isinstance(value, (datetime.time, datetime.date, datetime.datetime)): return value.isoformat() - elif isinstance(value, datetime.timedelta): - return str(value) - elif isinstance(value, uuid.UUID): + elif isinstance(value, (datetime.timedelta, uuid.UUID)): return str(value) else: return value def resolve_extracts( - extracts: Optional[Union[Dict[str, str], List[str], Tuple[str]]], + extracts: dict[str, str] | list[str] | tuple[str] | None, ) -> dict: if extracts is None: extracts = {} diff --git a/sqlite_utils/hookspecs.py b/sqlite_utils/hookspecs.py index a746619..73d1acc 100644 --- a/sqlite_utils/hookspecs.py +++ b/sqlite_utils/hookspecs.py @@ -1,8 +1,7 @@ import sqlite3 import click -from pluggy import HookimplMarker -from pluggy import HookspecMarker +from pluggy import HookimplMarker, HookspecMarker hookspec = HookspecMarker("sqlite_utils") hookimpl = HookimplMarker("sqlite_utils") diff --git a/sqlite_utils/migrations.py b/sqlite_utils/migrations.py index 00d0fa5..69397ba 100644 --- a/sqlite_utils/migrations.py +++ b/sqlite_utils/migrations.py @@ -1,19 +1,28 @@ -from collections.abc import Iterable -from dataclasses import dataclass import datetime -from typing import Callable, cast, TYPE_CHECKING +from collections.abc import Callable, Iterable +from dataclasses import dataclass +from typing import TYPE_CHECKING, Protocol, TypeVar, cast if TYPE_CHECKING: from sqlite_utils.db import Database, Table +class _MigrationFunction(Protocol): + __name__: str + + def __call__(self, db: "Database", /) -> None: ... + + +_MigrationFunctionT = TypeVar("_MigrationFunctionT", bound=_MigrationFunction) + + class Migrations: migrations_table = "_sqlite_migrations" @dataclass class _Migration: name: str - fn: Callable + fn: _MigrationFunction transactional: bool = True @dataclass @@ -32,7 +41,7 @@ class Migrations: def __call__( self, *, name: str | None = None, transactional: bool = True - ) -> Callable: + ) -> Callable[[_MigrationFunctionT], _MigrationFunctionT]: """ :param name: The name to use for this migration - if not provided, the name of the function will be used. @@ -43,13 +52,11 @@ class Migrations: example those that execute ``VACUUM``. """ - def inner(func: Callable) -> Callable: - migration_name = name or getattr(func, "__name__") + def inner(func: _MigrationFunctionT) -> _MigrationFunctionT: + migration_name = name or func.__name__ if any(m.name == migration_name for m in self._migrations): raise ValueError( - "Migration '{}' is already registered in set '{}'".format( - migration_name, self.name - ) + f"Migration '{migration_name}' is already registered in set '{self.name}'" ) self._migrations.append( self._Migration(migration_name, func, transactional) diff --git a/sqlite_utils/plugins.py b/sqlite_utils/plugins.py index 0aff7ff..10815b4 100644 --- a/sqlite_utils/plugins.py +++ b/sqlite_utils/plugins.py @@ -1,7 +1,7 @@ -from typing import Dict, List, Union +import sys import pluggy -import sys + from . import hookspecs pm: pluggy.PluginManager = pluggy.PluginManager("sqlite_utils") @@ -17,13 +17,13 @@ def ensure_plugins_loaded() -> None: _plugins_loaded = True -def get_plugins() -> List[Dict[str, Union[str, List[str]]]]: +def get_plugins() -> list[dict[str, str | list[str]]]: ensure_plugins_loaded() - plugins: List[Dict[str, Union[str, List[str]]]] = [] + plugins: list[dict[str, str | list[str]]] = [] plugin_to_distinfo = dict(pm.list_plugin_distinfo()) for plugin in pm.get_plugins(): hookcallers = pm.get_hookcallers(plugin) or [] - plugin_info: Dict[str, Union[str, List[str]]] = { + plugin_info: dict[str, str | list[str]] = { "name": plugin.__name__, "hooks": [h.name for h in hookcallers], } diff --git a/sqlite_utils/recipes.py b/sqlite_utils/recipes.py index 55b55a4..d28a099 100644 --- a/sqlite_utils/recipes.py +++ b/sqlite_utils/recipes.py @@ -1,9 +1,9 @@ from __future__ import annotations -from typing import Callable, Optional +import json +from collections.abc import Callable from dateutil import parser -import json IGNORE: object = object() SET_NULL: object = object() @@ -13,8 +13,8 @@ def parsedate( value: str, dayfirst: bool = False, yearfirst: bool = False, - errors: Optional[object] = None, -) -> Optional[str]: + errors: object | None = None, +) -> str | None: """ Parse a date and convert it to ISO date format: yyyy-mm-dd \b @@ -44,8 +44,8 @@ def parsedatetime( value: str, dayfirst: bool = False, yearfirst: bool = False, - errors: Optional[object] = None, -) -> Optional[str]: + errors: object | None = None, +) -> str | None: """ Parse a datetime and convert it to ISO datetime format: yyyy-mm-ddTHH:MM:SS \b diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index b39b117..ed5a558 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -9,20 +9,11 @@ import itertools import json import os import sys +from collections.abc import Callable, Generator, Iterable, Iterator from typing import ( + TYPE_CHECKING, Any, BinaryIO, - Callable, - Dict, - Generator, - Iterable, - Iterator, - List, - Optional, - Set, - Tuple, - Type, - TYPE_CHECKING, TypeVar, Union, cast, @@ -33,8 +24,8 @@ import click from . import recipes if TYPE_CHECKING: - import sqlite3 # noqa: F401 - from sqlite3 import dbapi2 # noqa: F401 + import sqlite3 + from sqlite3 import dbapi2 OperationalError = dbapi2.OperationalError else: @@ -44,7 +35,7 @@ else: OperationalError = dbapi2.OperationalError except ImportError: import sqlite3 # noqa: F401 - from sqlite3 import dbapi2 # noqa: F401 + from sqlite3 import dbapi2 OperationalError = dbapi2.OperationalError @@ -61,8 +52,8 @@ 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, List[str]] -Row = Dict[str, RowValue] +RowValue = None | int | float | str | bytes | bool | list[str] +Row = dict[str, RowValue] T = TypeVar("T") @@ -103,7 +94,7 @@ def maximize_csv_field_size_limit() -> None: field_size_limit = int(field_size_limit / 10) -def find_spatialite() -> Optional[str]: +def find_spatialite() -> str | None: """ The ``find_spatialite()`` function searches for the `SpatiaLite `__ SQLite extension in some common places. It returns a string path to the location, or ``None`` if SpatiaLite was not found. @@ -132,9 +123,9 @@ def find_spatialite() -> Optional[str]: def suggest_column_types( - records: Iterable[Dict[str, Any]], -) -> Dict[str, type]: - all_column_types: Dict[str, Set[type]] = {} + 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)) @@ -142,9 +133,9 @@ def suggest_column_types( def types_for_column_types( - all_column_types: Dict[str, Set[type]], -) -> Dict[str, type]: - column_types: Dict[str, type] = {} + 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: @@ -153,7 +144,7 @@ def types_for_column_types( if {None.__class__} == types: t = str elif len(types) == 1: - t = list(types)[0] + t = next(iter(types)) # But if it's a subclass of list / tuple / dict, use str # instead as we will be storing it as JSON in the table for superclass in (list, tuple, dict): @@ -190,7 +181,7 @@ def column_affinity(column_type: str) -> type: return float -def decode_base64_values(doc: Dict[str, Any]) -> Dict[str, Any]: +def decode_base64_values(doc: dict[str, Any]) -> dict[str, Any]: # Looks for '{"$base64": true..., "encoded": ...}' values and decodes them to_fix = [ k @@ -263,9 +254,9 @@ class RowError(Exception): def _extra_key_strategy( - reader: Iterable[Dict[Optional[str], object]], - ignore_extras: Optional[bool] = False, - extras_key: Optional[str] = None, + reader: Iterable[dict[str | None, object]], + ignore_extras: bool | None = False, + extras_key: str | None = None, ) -> Iterable[Row]: # Logic for handling CSV rows with more values than there are headings for row in reader: @@ -279,9 +270,7 @@ def _extra_key_strategy( yield cast(Row, row) elif not extras_key: extras = row.pop(None) - raise RowError( - "Row {} contained these extra values: {}".format(row, extras) - ) + raise RowError(f"Row {row} contained these extra values: {extras}") else: extras_value = row.pop(None) row_out = cast(Row, row) @@ -291,12 +280,12 @@ def _extra_key_strategy( def rows_from_file( fp: BinaryIO, - format: Optional[Format] = None, - dialect: Optional[Type[csv.Dialect]] = None, - encoding: Optional[str] = None, - ignore_extras: Optional[bool] = False, - extras_key: Optional[str] = None, -) -> Tuple[Iterable[Row], Format]: + format: Format | None = None, + dialect: type[csv.Dialect] | None = None, + encoding: str | None = None, + ignore_extras: bool | None = False, + extras_key: str | None = None, +) -> tuple[Iterable[Row], Format]: """ Load a sequence of dictionaries from a file-like object containing one of four different formats. @@ -363,7 +352,7 @@ def rows_from_file( ) return ( _extra_key_strategy( - cast(Iterable[Dict[Optional[str], object]], rows), + cast(Iterable[dict[str | None, object]], rows), ignore_extras, extras_key, ), @@ -379,7 +368,7 @@ def rows_from_file( raise TypeError( "rows_from_file() requires a file-like object that supports peek(), such as io.BytesIO" ) - if first_bytes.startswith(b"[") or first_bytes.startswith(b"{"): + if first_bytes.startswith((b"[", b"{")): # TODO: Detect newline-JSON return rows_from_file(buffered, format=Format.JSON) else: @@ -393,7 +382,7 @@ def rows_from_file( detected_format = Format.TSV if dialect.delimiter == "\t" else Format.CSV return ( _extra_key_strategy( - cast(Iterable[Dict[Optional[str], object]], rows), + cast(Iterable[dict[str | None, object]], rows), ignore_extras, extras_key, ), @@ -425,9 +414,9 @@ class TypeTracker: """ def __init__(self) -> None: - self.trackers: Dict[str, "ValueTracker"] = {} + self.trackers: dict[str, ValueTracker] = {} - def wrap(self, iterator: Iterable[Dict[str, Any]]) -> Iterable[Dict[str, Any]]: + 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. @@ -441,7 +430,7 @@ class TypeTracker: yield row @property - def types(self) -> Dict[str, str]: + def types(self) -> dict[str, str]: """ A dictionary mapping column names to their detected types. This can be passed to the ``db[table_name].transform(types=tracker.types)`` method. @@ -450,17 +439,15 @@ class TypeTracker: class ValueTracker: - couldbe: Dict[str, Callable[[object], bool]] + couldbe: dict[str, Callable[[object], bool]] def __init__(self) -> None: self.couldbe = {key: getattr(self, "test_" + key) for key in self.get_tests()} @classmethod - def get_tests(cls) -> List[str]: + def get_tests(cls) -> list[str]: return [ - key.split("test_")[-1] - for key in cls.__dict__.keys() - if key.startswith("test_") + key.split("test_")[-1] for key in cls.__dict__ if key.startswith("test_") ] def test_integer(self, value: object) -> bool: @@ -492,7 +479,7 @@ class ValueTracker: def evaluate(self, value: object) -> None: if not value or not self.couldbe: return - not_these: List[str] = [] + not_these: list[str] = [] for name, test in self.couldbe.items(): if not test(value): not_these.append(name) @@ -524,14 +511,14 @@ def progressbar(*args: Iterable[T], **kwargs: Any) -> Generator[Any, None, None] def _compile_code( code: str, imports: Iterable[str], variable: str = "value" ) -> Callable[..., Any]: - globals_dict: Dict[str, Any] = {"r": recipes, "recipes": recipes} + globals_dict: dict[str, Any] = {"r": recipes, "recipes": recipes} # Handle imports first so they're available for all approaches for import_ in imports: globals_dict[import_.split(".")[0]] = __import__(import_) # If user defined a convert() function, return that try: - exec(code, globals_dict) + exec(code, globals_dict) # noqa: S102 return cast(Callable[..., object], globals_dict["convert"]) except (AttributeError, SyntaxError, NameError, KeyError, TypeError): pass @@ -542,20 +529,20 @@ def _compile_code( fn = eval(code, globals_dict) if callable(fn): return cast(Callable[..., object], fn) - except Exception: + except Exception: # noqa: BLE001, S110 pass # Try compiling their code as a function instead body_variants = [code] # If single line and no 'return', try adding the return if "\n" not in code and not code.strip().startswith("return "): - body_variants.insert(0, "return {}".format(code)) + body_variants.insert(0, f"return {code}") code_o = None for variant in body_variants: - new_code = ["def fn({}):".format(variable)] + new_code = [f"def fn({variable}):"] for line in variant.split("\n"): - new_code.append(" {}".format(line)) + new_code.append(f" {line}") try: code_o = compile("\n".join(new_code), "", "exec") break @@ -566,7 +553,7 @@ def _compile_code( if code_o is None: raise SyntaxError("Could not compile code") - exec(code_o, globals_dict) + exec(code_o, globals_dict) # noqa: S102 return cast(Callable[..., object], globals_dict["fn"]) @@ -582,7 +569,7 @@ def chunks(sequence: Iterable[T], size: int) -> Iterable[Iterable[T]]: yield itertools.chain([item], itertools.islice(iterator, size - 1)) -def hash_record(record: Dict[str, Any], keys: Optional[Iterable[str]] = None) -> str: +def hash_record(record: dict[str, Any], keys: Iterable[str] | None = None) -> str: """ ``record`` should be a Python dictionary. Returns a sha1 hash of the keys and values in that record. @@ -603,7 +590,7 @@ def hash_record(record: Dict[str, Any], 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: Dict[str, Any] = record + to_hash: dict[str, Any] = record if keys is not None: to_hash = {key: record[key] for key in keys} return hashlib.sha1( @@ -613,7 +600,7 @@ def hash_record(record: Dict[str, Any], keys: Optional[Iterable[str]] = None) -> ).hexdigest() -def dedupe_keys(keys: Iterable[str]) -> List[str]: +def dedupe_keys(keys: Iterable[str]) -> list[str]: """ Rename duplicates in a list of column names so every name is unique, by appending ``_2``, ``_3``... to later occurrences - skipping any @@ -636,7 +623,7 @@ def dedupe_keys(keys: Iterable[str]) -> List[str]: new_key = key suffix = 2 while new_key in seen or new_key in taken: - new_key = "{}_{}".format(key, suffix) + new_key = f"{key}_{suffix}" suffix += 1 key = new_key seen.add(key) @@ -644,7 +631,7 @@ def dedupe_keys(keys: Iterable[str]) -> List[str]: return result -def _flatten(d: Dict[str, Any]) -> Generator[Tuple[str, Any], None, None]: +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): @@ -653,7 +640,7 @@ def _flatten(d: Dict[str, Any]) -> Generator[Tuple[str, Any], None, None]: yield key, value -def flatten(row: Dict[str, Any]) -> Dict[str, Any]: +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}`` diff --git a/tests/conftest.py b/tests/conftest.py index 728db7b..a4eb860 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,6 +1,7 @@ +import pytest + from sqlite_utils import Database from sqlite_utils.utils import sqlite3 -import pytest CREATE_TABLES = """ create table Gosh (c1 text, c2 text, c3 text); @@ -55,7 +56,7 @@ def close_all_databases(): for db in databases: try: db.close() - except Exception: + except sqlite3.Error: pass diff --git a/tests/test_analyze_tables.py b/tests/test_analyze_tables.py index a2ce585..a51bba6 100644 --- a/tests/test_analyze_tables.py +++ b/tests/test_analyze_tables.py @@ -1,9 +1,11 @@ -from sqlite_utils.db import Database, ColumnDetails -from sqlite_utils import cli -from click.testing import CliRunner -import pytest import sqlite3 +import pytest +from click.testing import CliRunner + +from sqlite_utils import cli +from sqlite_utils.db import ColumnDetails, Database + @pytest.fixture def db_to_analyze(fresh_db): diff --git a/tests/test_atomic.py b/tests/test_atomic.py index c3fd02f..ba16ca5 100644 --- a/tests/test_atomic.py +++ b/tests/test_atomic.py @@ -28,11 +28,13 @@ from sqlite_utils.utils import sqlite3 END; """, [ - "CREATE TRIGGER t_ai AFTER INSERT ON t\n" - " BEGIN\n" - " UPDATE t SET value = 'a;b' WHERE id = new.id;\n" - " INSERT INTO log VALUES ('x;y');\n" - " END;" + ( + "CREATE TRIGGER t_ai AFTER INSERT ON t\n" + " BEGIN\n" + " UPDATE t SET value = 'a;b' WHERE id = new.id;\n" + " INSERT INTO log VALUES ('x;y');\n" + " END;" + ) ], ), ), @@ -49,10 +51,9 @@ def test_atomic_commits(fresh_db): def test_atomic_rolls_back(fresh_db): - with pytest.raises(RuntimeError): - with fresh_db.atomic(): - fresh_db["dogs"].insert({"id": 1, "name": "Cleo"}, pk="id") - raise RuntimeError("boom") + with pytest.raises(RuntimeError), fresh_db.atomic(): + fresh_db["dogs"].insert({"id": 1, "name": "Cleo"}, pk="id") + raise RuntimeError("boom") assert not fresh_db["dogs"].exists() @@ -62,10 +63,9 @@ def test_nested_atomic_rolls_back_to_savepoint(fresh_db): with fresh_db.atomic(): fresh_db["dogs"].insert({"id": 1, "name": "Cleo"}) - with pytest.raises(RuntimeError): - with fresh_db.atomic(): - fresh_db["dogs"].insert({"id": 2, "name": "Pancakes"}) - raise RuntimeError("boom") + with pytest.raises(RuntimeError), fresh_db.atomic(): + fresh_db["dogs"].insert({"id": 2, "name": "Pancakes"}) + raise RuntimeError("boom") fresh_db["dogs"].insert({"id": 3, "name": "Marnie"}) assert list(fresh_db["dogs"].rows) == [ @@ -75,20 +75,18 @@ def test_nested_atomic_rolls_back_to_savepoint(fresh_db): def test_outer_atomic_rolls_back_released_savepoint(fresh_db): - with pytest.raises(RuntimeError): + with pytest.raises(RuntimeError), fresh_db.atomic(): + fresh_db["dogs"].insert({"id": 1, "name": "Cleo"}, pk="id") with fresh_db.atomic(): - fresh_db["dogs"].insert({"id": 1, "name": "Cleo"}, pk="id") - with fresh_db.atomic(): - fresh_db["dogs"].insert({"id": 2, "name": "Pancakes"}) - raise RuntimeError("boom") + fresh_db["dogs"].insert({"id": 2, "name": "Pancakes"}) + raise RuntimeError("boom") assert not fresh_db["dogs"].exists() def test_executescript_does_not_commit_open_atomic_block(fresh_db): - with pytest.raises(RuntimeError): - with fresh_db.atomic(): - fresh_db.executescript(""" + with pytest.raises(RuntimeError), fresh_db.atomic(): + fresh_db.executescript(""" CREATE TABLE dogs(id INTEGER PRIMARY KEY, name TEXT); CREATE TRIGGER dogs_ai AFTER INSERT ON dogs BEGIN @@ -97,7 +95,7 @@ def test_executescript_does_not_commit_open_atomic_block(fresh_db): -- This comment has a semicolon; INSERT INTO dogs VALUES (1, 'Cleo; the first'); """) - raise RuntimeError("boom") + raise RuntimeError("boom") assert not fresh_db["dogs"].exists() @@ -105,11 +103,10 @@ def test_executescript_does_not_commit_open_atomic_block(fresh_db): def test_transform_does_not_commit_open_atomic_block(fresh_db): fresh_db["dogs"].insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id") - with pytest.raises(RuntimeError): - with fresh_db.atomic(): - fresh_db["dogs"].insert({"id": 2, "name": "Pancakes", "age": "6"}) - fresh_db["dogs"].transform(rename={"age": "dog_age"}) - raise RuntimeError("boom") + with pytest.raises(RuntimeError), fresh_db.atomic(): + fresh_db["dogs"].insert({"id": 2, "name": "Pancakes", "age": "6"}) + fresh_db["dogs"].transform(rename={"age": "dog_age"}) + raise RuntimeError("boom") assert ( fresh_db["dogs"].schema @@ -149,10 +146,9 @@ def test_transform_parent_table_with_foreign_keys_rolls_back(fresh_db): foreign_keys={"author_id"}, ) - with pytest.raises(RuntimeError): - with fresh_db.atomic(): - fresh_db["authors"].transform(rename={"name": "full_name"}) - raise RuntimeError("boom") + with pytest.raises(RuntimeError), fresh_db.atomic(): + fresh_db["authors"].transform(rename={"name": "full_name"}) + raise RuntimeError("boom") assert ( fresh_db["authors"].schema @@ -354,9 +350,11 @@ def test_atomic_preserves_error_from_transaction_destroying_trigger(fresh_db): # with "cannot rollback - no transaction is active" fresh_db.execute("create table t (id integer primary key, v text)") fresh_db.execute(TRIGGER_SQL) - with pytest.raises(sqlite3.IntegrityError, match="trigger says no"): - with fresh_db.atomic(): - fresh_db.execute("insert into t (v) values ('bad')") + with ( + pytest.raises(sqlite3.IntegrityError, match="trigger says no"), + fresh_db.atomic(), + ): + fresh_db.execute("insert into t (v) values ('bad')") assert not fresh_db.conn.in_transaction @@ -367,16 +365,17 @@ def test_nested_atomic_preserves_error_from_transaction_destroying_trigger( # "no such savepoint" from ROLLBACK TO SAVEPOINT fresh_db.execute("create table t (id integer primary key, v text)") fresh_db.execute(TRIGGER_SQL) - with pytest.raises(sqlite3.IntegrityError, match="trigger says no"): - with fresh_db.atomic(): - with fresh_db.atomic(): - fresh_db.execute("insert into t (v) values ('bad')") + with ( + pytest.raises(sqlite3.IntegrityError, match="trigger says no"), + fresh_db.atomic(), + fresh_db.atomic(), + ): + fresh_db.execute("insert into t (v) values ('bad')") assert not fresh_db.conn.in_transaction def test_atomic_preserves_error_from_insert_or_rollback(fresh_db): fresh_db["t"].insert({"id": 1}, pk="id") - with pytest.raises(sqlite3.IntegrityError): - with fresh_db.atomic(): - fresh_db.execute("insert or rollback into t (id) values (1)") + with pytest.raises(sqlite3.IntegrityError), fresh_db.atomic(): + fresh_db.execute("insert or rollback into t (id) values (1)") assert not fresh_db.conn.in_transaction diff --git a/tests/test_cli.py b/tests/test_cli.py index a2135b0..a1e072f 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1,14 +1,16 @@ -from sqlite_utils import cli, Database -from sqlite_utils.db import Index, ForeignKey -from click.testing import CliRunner -from pathlib import Path -import subprocess -import sqlite3 -import sys import json import os -import pytest +import sqlite3 +import subprocess +import sys import textwrap +from pathlib import Path + +import pytest +from click.testing import CliRunner + +from sqlite_utils import Database, cli +from sqlite_utils.db import ForeignKey, Index def write_json(file_path, data): @@ -21,7 +23,7 @@ def _supports_pragma_function_list(): try: db.execute("select * from pragma_function_list()") return True - except Exception: + except sqlite3.DatabaseError: return False finally: db.close() @@ -184,9 +186,9 @@ def test_output_table(db_path, options, expected): db["rows"].insert_all( [ { - "c1": "verb{}".format(i), - "c2": "noun{}".format(i), - "c3": "adjective{}".format(i), + "c1": f"verb{i}", + "c2": f"noun{i}", + "c3": f"adjective{i}", } for i in range(4) ] @@ -678,9 +680,9 @@ def test_optimize(db_path, tables): db[table].insert_all( [ { - "c1": "verb{}".format(i), - "c2": "noun{}".format(i), - "c3": "adjective{}".format(i), + "c1": f"verb{i}", + "c2": f"noun{i}", + "c3": f"adjective{i}", } for i in range(10000) ] @@ -704,9 +706,9 @@ def test_rebuild_fts_fixes_docsize_error(db_path): db = Database(db_path, recursive_triggers=False) records = [ { - "c1": "verb{}".format(i), - "c2": "noun{}".format(i), - "c3": "adjective{}".format(i), + "c1": f"verb{i}", + "c2": f"noun{i}", + "c3": f"adjective{i}", } for i in range(10000) ] @@ -1019,16 +1021,14 @@ def test_query_json_binary(db_path): "data": { "$base64": True, "encoded": ( - ( - "eJzt0c1xAyEMBeC7q1ABHleR3HxNAQrIjmb4M0gelx+RTY7p4N2WBYT0vmufUknH" - "8kq5lz5pqRFXsTOl3pYkE/NJnHXoStruJEVjc0mOCyTqq/ZMJnXEZW1Js2ZvRm5U+" - "DPKk9hRWqjyvTFx0YfzhT6MpGmN2lR1fzxjyfVMD9dFrS+bnkleMpMam/ZGXgrX1I" - "/K+5Au3S/9lNQRh0k4Gq/RUz8GiKfsQm+7JLsJ6fTo5JhVG00ZU76kZZkxePx49uI" - "jnpNoJyYlWUsoaSl/CcVATje/Kxu13RANnrHweaH3V5Jh4jvGyKCnxJLiXPKhmW3f" - "iCnG7Jql7RR3UvFo8jJ4z039dtOkTFmWzL1be9lt8A5II471m6vXy+l0BR/4wAc+8" - "IEPfOADH/jABz7wgQ984AMf+MAHPvCBD3zgAx/4wAc+8IEPfOADH/jABz7wgQ984A" - "Mf+MAHPvCBD3zgAx/4wAc+8IEPfOADH/jABz7wgQ984PuP7xubBoN9" - ) + "eJzt0c1xAyEMBeC7q1ABHleR3HxNAQrIjmb4M0gelx+RTY7p4N2WBYT0vmufUknH" + "8kq5lz5pqRFXsTOl3pYkE/NJnHXoStruJEVjc0mOCyTqq/ZMJnXEZW1Js2ZvRm5U+" + "DPKk9hRWqjyvTFx0YfzhT6MpGmN2lR1fzxjyfVMD9dFrS+bnkleMpMam/ZGXgrX1I" + "/K+5Au3S/9lNQRh0k4Gq/RUz8GiKfsQm+7JLsJ6fTo5JhVG00ZU76kZZkxePx49uI" + "jnpNoJyYlWUsoaSl/CcVATje/Kxu13RANnrHweaH3V5Jh4jvGyKCnxJLiXPKhmW3f" + "iCnG7Jql7RR3UvFo8jJ4z039dtOkTFmWzL1be9lt8A5II471m6vXy+l0BR/4wAc+8" + "IEPfOADH/jABz7wgQ984AMf+MAHPvCBD3zgAx/4wAc+8IEPfOADH/jABz7wgQ984A" + "Mf+MAHPvCBD3zgAx/4wAc+8IEPfOADH/jABz7wgQ984PuP7xubBoN9" ), }, } @@ -2114,11 +2114,13 @@ _common_other_schema = ( ), ( ["--rename", "name", "name2"], - 'CREATE TABLE "trees" (\n' - ' "id" INTEGER PRIMARY KEY,\n' - ' "address" TEXT,\n' - ' "species_id" INTEGER REFERENCES "species"("id")\n' - ")", + ( + 'CREATE TABLE "trees" (\n' + ' "id" INTEGER PRIMARY KEY,\n' + ' "address" TEXT,\n' + ' "species_id" INTEGER REFERENCES "species"("id")\n' + ")" + ), 'CREATE TABLE "species" (\n "id" INTEGER PRIMARY KEY,\n "species" TEXT\n)', ), ], @@ -2137,9 +2139,9 @@ def test_extract(db_path, args, expected_table_schema, expected_other_schema): assert result.exit_code == 0 schema = db["trees"].schema assert schema == expected_table_schema - other_schema = [t for t in db.tables if t.name not in ("trees", "Gosh", "Gosh2")][ - 0 - ].schema + other_schema = next( + t for t in db.tables if t.name not in ("trees", "Gosh", "Gosh2") + ).schema assert other_schema == expected_other_schema @@ -2431,7 +2433,7 @@ def test_long_csv_column_value(tmpdir): with open(csv_path, "w") as csv_file: long_string = "a" * 131073 csv_file.write("id,text\n") - csv_file.write("1,{}\n".format(long_string)) + csv_file.write(f"1,{long_string}\n") result = CliRunner().invoke( cli.cli, ["insert", db_path, "bigtable", csv_path, "--csv"], @@ -2457,8 +2459,8 @@ def test_import_no_headers(tmpdir, args, tsv): csv_path = str(tmpdir / "test.csv") with open(csv_path, "w") as csv_file: sep = "\t" if tsv else "," - csv_file.write("Cleo{sep}Dog{sep}5\n".format(sep=sep)) - csv_file.write("Tracy{sep}Spider{sep}7\n".format(sep=sep)) + csv_file.write(f"Cleo{sep}Dog{sep}5\n") + csv_file.write(f"Tracy{sep}Spider{sep}7\n") result = CliRunner().invoke( cli.cli, ["insert", db_path, "creatures", csv_path] + args + ["--no-detect-types"], @@ -2690,7 +2692,9 @@ def test_integer_overflow_error(tmpdir): def test_python_dash_m(): "Tool can be run using python -m sqlite_utils" result = subprocess.run( - [sys.executable, "-m", "sqlite_utils", "--help"], stdout=subprocess.PIPE + [sys.executable, "-m", "sqlite_utils", "--help"], + stdout=subprocess.PIPE, + check=False, ) assert result.returncode == 0 assert b"Commands for interacting with a SQLite database" in result.stdout @@ -2830,14 +2834,14 @@ def test_load_extension(entrypoint, should_pass, should_fail): for func in should_pass: result = CliRunner().invoke( cli.cli, - ["memory", "select {}()".format(func), "--load-extension", ext], + ["memory", f"select {func}()", "--load-extension", ext], catch_exceptions=False, ) assert result.exit_code == 0 for func in should_fail: result = CliRunner().invoke( cli.cli, - ["memory", "select {}()".format(func), "--load-extension", ext], + ["memory", f"select {func}()", "--load-extension", ext], catch_exceptions=False, ) assert result.exit_code == 1 diff --git a/tests/test_cli_bulk.py b/tests/test_cli_bulk.py index 514f4ac..932269b 100644 --- a/tests/test_cli_bulk.py +++ b/tests/test_cli_bulk.py @@ -1,11 +1,13 @@ -from click.testing import CliRunner -from sqlite_utils import cli, Database import pathlib -import pytest import subprocess import sys import time +import pytest +from click.testing import CliRunner + +from sqlite_utils import Database, cli + @pytest.fixture def test_db_and_path(tmpdir): diff --git a/tests/test_cli_convert.py b/tests/test_cli_convert.py index 6c3f5c5..65543b1 100644 --- a/tests/test_cli_convert.py +++ b/tests/test_cli_convert.py @@ -1,10 +1,12 @@ -from click.testing import CliRunner -from sqlite_utils import cli -import sqlite_utils import json -import textwrap import pathlib +import textwrap + import pytest +from click.testing import CliRunner + +import sqlite_utils +from sqlite_utils import cli @pytest.fixture @@ -50,7 +52,7 @@ def test_convert_code(fresh_db_and_path, code): cli.cli, ["convert", db_path, "t", "text", code], catch_exceptions=False ) assert result.exit_code == 0, result.output - value = list(db["t"].rows)[0]["text"] + value = next(iter(db["t"].rows))["text"] assert value == "Spooktober" @@ -442,7 +444,7 @@ def test_recipe_jsonsplit(tmpdir, delimiter): ) code = "r.jsonsplit(value)" if delimiter: - code = 'recipes.jsonsplit(value, delimiter="{}")'.format(delimiter) + code = f'recipes.jsonsplit(value, delimiter="{delimiter}")' args = ["convert", db_path, "example", "tags", code] result = CliRunner().invoke(cli.cli, args) assert result.exit_code == 0, result.output @@ -470,7 +472,7 @@ def test_recipe_jsonsplit_type(fresh_db_and_path, type, expected_array): ) code = "r.jsonsplit(value)" if type: - code = "recipes.jsonsplit(value, type={})".format(type) + code = f"recipes.jsonsplit(value, type={type})" args = ["convert", db_path, "example", "records", code] result = CliRunner().invoke(cli.cli, args) assert result.exit_code == 0, result.output diff --git a/tests/test_cli_insert.py b/tests/test_cli_insert.py index df6f80c..eefb3fa 100644 --- a/tests/test_cli_insert.py +++ b/tests/test_cli_insert.py @@ -1,11 +1,13 @@ -from sqlite_utils import cli, Database -from click.testing import CliRunner import json -import pytest import subprocess import sys import time +import pytest +from click.testing import CliRunner + +from sqlite_utils import Database, cli + def test_insert_simple(tmpdir): json_path = str(tmpdir / "dog.json") @@ -99,7 +101,7 @@ def test_insert_with_primary_keys(db_path, tmpdir, args, expected_pks): def test_insert_multiple_with_primary_key(db_path, tmpdir): json_path = str(tmpdir / "dogs.json") - dogs = [{"id": i, "name": "Cleo {}".format(i), "age": i + 3} for i in range(1, 21)] + dogs = [{"id": i, "name": f"Cleo {i}", "age": i + 3} for i in range(1, 21)] with open(json_path, "w") as fp: fp.write(json.dumps(dogs)) result = CliRunner().invoke( @@ -114,7 +116,7 @@ def test_insert_multiple_with_primary_key(db_path, tmpdir): def test_insert_multiple_with_compound_primary_key(db_path, tmpdir): json_path = str(tmpdir / "dogs.json") dogs = [ - {"breed": "mixed", "id": i, "name": "Cleo {}".format(i), "age": i + 3} + {"breed": "mixed", "id": i, "name": f"Cleo {i}", "age": i + 3} for i in range(1, 21) ] with open(json_path, "w") as fp: @@ -140,8 +142,7 @@ def test_insert_multiple_with_compound_primary_key(db_path, tmpdir): def test_insert_not_null_default(db_path, tmpdir): json_path = str(tmpdir / "dogs.json") dogs = [ - {"id": i, "name": "Cleo {}".format(i), "age": i + 3, "score": 10} - for i in range(1, 21) + {"id": i, "name": f"Cleo {i}", "age": i + 3, "score": 10} for i in range(1, 21) ] with open(json_path, "w") as fp: fp.write(json.dumps(dogs)) @@ -587,7 +588,7 @@ def test_insert_streaming_batch_size_1(db_path): return tries += 1 if tries > 10: - assert False, "Expected {}, got {}".format(expected, rows) + assert False, f"Expected {expected}, got {rows}" time.sleep(tries * 0.1) try_until([{"name": "Azi"}]) diff --git a/tests/test_cli_memory.py b/tests/test_cli_memory.py index 2ed4aaa..4fb4fb3 100644 --- a/tests/test_cli_memory.py +++ b/tests/test_cli_memory.py @@ -1,5 +1,6 @@ -import click import json + +import click import pytest from click.testing import CliRunner @@ -28,7 +29,7 @@ def test_memory_csv(tmpdir, sql_from, use_stdin): fp.write(content) result = CliRunner().invoke( cli.cli, - ["memory", csv_path, "select * from {}".format(sql_from), "--nl"], + ["memory", csv_path, f"select * from {sql_from}", "--nl"], input=input, ) assert result.exit_code == 0 @@ -53,7 +54,7 @@ def test_memory_tsv(tmpdir, use_stdin): sql_from = "chickens" result = CliRunner().invoke( cli.cli, - ["memory", path, "select * from {}".format(sql_from)], + ["memory", path, f"select * from {sql_from}"], input=input, ) assert result.exit_code == 0, result.output @@ -79,7 +80,7 @@ def test_memory_json(tmpdir, use_stdin): sql_from = "chickens" result = CliRunner().invoke( cli.cli, - ["memory", path, "select * from {}".format(sql_from)], + ["memory", path, f"select * from {sql_from}"], input=input, ) assert result.exit_code == 0, result.output @@ -105,7 +106,7 @@ def test_memory_json_nl(tmpdir, use_stdin): sql_from = "chickens" result = CliRunner().invoke( cli.cli, - ["memory", path, "select * from {}".format(sql_from)], + ["memory", path, f"select * from {sql_from}"], input=input, ) assert result.exit_code == 0, result.output @@ -135,7 +136,7 @@ def test_memory_csv_encoding(tmpdir, use_stdin): CliRunner() .invoke( cli.cli, - ["memory", csv_path, "select * from {}".format(sql_from), "--nl"], + ["memory", csv_path, f"select * from {sql_from}", "--nl"], input=input, ) .exit_code diff --git a/tests/test_cli_migrate.py b/tests/test_cli_migrate.py index 0f29e36..f49ef10 100644 --- a/tests/test_cli_migrate.py +++ b/tests/test_cli_migrate.py @@ -1,7 +1,8 @@ import pathlib -from click.testing import CliRunner import pytest +from click.testing import CliRunner + import sqlite_utils import sqlite_utils.cli diff --git a/tests/test_column_affinity.py b/tests/test_column_affinity.py index fb8f340..fa23345 100644 --- a/tests/test_column_affinity.py +++ b/tests/test_column_affinity.py @@ -1,4 +1,5 @@ import pytest + from sqlite_utils.utils import column_affinity EXAMPLES = [ @@ -41,5 +42,5 @@ def test_column_affinity(column_def, expected_type): @pytest.mark.parametrize("column_def,expected_type", EXAMPLES) def test_columns_dict(fresh_db, column_def, expected_type): - fresh_db.execute("create table foo (col {})".format(column_def)) + fresh_db.execute(f"create table foo (col {column_def})") assert {"col": expected_type} == fresh_db["foo"].columns_dict diff --git a/tests/test_constructor.py b/tests/test_constructor.py index a619fba..4282969 100644 --- a/tests/test_constructor.py +++ b/tests/test_constructor.py @@ -1,8 +1,10 @@ +import sys + +import pytest + from sqlite_utils import Database from sqlite_utils.db import TransactionError from sqlite_utils.utils import sqlite3 -import pytest -import sys def test_recursive_triggers(): diff --git a/tests/test_convert.py b/tests/test_convert.py index ea3fd96..879267a 100644 --- a/tests/test_convert.py +++ b/tests/test_convert.py @@ -1,6 +1,7 @@ -from sqlite_utils.db import BadMultiValues import pytest +from sqlite_utils.db import BadMultiValues + @pytest.mark.parametrize( "columns,fn,expected", diff --git a/tests/test_create.py b/tests/test_create.py index d281eb4..40746bf 100644 --- a/tests/test_create.py +++ b/tests/test_create.py @@ -1,26 +1,28 @@ -from sqlite_utils.db import ( - Index, - Database, - DescIndex, - AlterError, - InvalidColumns, - NoObviousTable, - OperationalError, - ForeignKey, - Table, - View, - NoTable, - NoView, -) -from sqlite_utils.utils import hash_record, sqlite3 import collections import datetime import decimal import json import pathlib -import pytest import uuid +import pytest + +from sqlite_utils.db import ( + AlterError, + Database, + DescIndex, + ForeignKey, + Index, + InvalidColumns, + NoObviousTable, + NoTable, + NoView, + OperationalError, + Table, + View, +) +from sqlite_utils.utils import hash_record, sqlite3 + try: import pandas as pd # type: ignore except ImportError: @@ -699,7 +701,7 @@ def test_bulk_insert_more_than_999_values(fresh_db): "num_columns,should_error", ((900, False), (999, False), (1000, True)) ) def test_error_if_more_than_999_columns(fresh_db, num_columns, should_error): - record = dict([("c{}".format(i), i) for i in range(num_columns)]) + record = {f"c{i}": i for i in range(num_columns)} if should_error: with pytest.raises(ValueError): fresh_db["big"].insert(record) @@ -718,17 +720,9 @@ def test_columns_not_in_first_record_should_not_cause_batch_to_be_too_large(fres records = [ {"c0": "first record"}, # one column in first record -> batch size = 999 # fill out the batch with 99 records with enough columns to exceed THRESHOLD - *[ - dict([("c{}".format(i), j) for i in range(extra_columns)]) - for j in range(batch_size - 1) - ], + *[{f"c{i}": j for i in range(extra_columns)} for j in range(batch_size - 1)], ] - try: - fresh_db["too_many_columns"].insert_all( - records, alter=True, batch_size=batch_size - ) - except sqlite3.OperationalError: - raise + fresh_db["too_many_columns"].insert_all(records, alter=True, batch_size=batch_size) @pytest.mark.parametrize( @@ -910,7 +904,7 @@ def test_insert_list_nested_unicode(fresh_db): def test_insert_uuid(fresh_db): uuid4 = uuid.uuid4() fresh_db["test"].insert({"uuid": uuid4}) - row = list(fresh_db["test"].rows)[0] + row = next(iter(fresh_db["test"].rows)) assert {"uuid"} == row.keys() assert isinstance(row["uuid"], str) assert row["uuid"] == str(uuid4) @@ -918,16 +912,14 @@ def test_insert_uuid(fresh_db): def test_insert_memoryview(fresh_db): fresh_db["test"].insert({"data": memoryview(b"hello")}) - row = list(fresh_db["test"].rows)[0] + row = next(iter(fresh_db["test"].rows)) assert {"data"} == row.keys() assert isinstance(row["data"], bytes) assert row["data"] == b"hello" def test_insert_thousands_using_generator(fresh_db): - fresh_db["test"].insert_all( - {"i": i, "word": "word_{}".format(i)} for i in range(10000) - ) + fresh_db["test"].insert_all({"i": i, "word": f"word_{i}"} for i in range(10000)) assert [{"name": "i", "type": "INTEGER"}, {"name": "word", "type": "TEXT"}] == [ {"name": col.name, "type": col.type} for col in fresh_db["test"].columns ] @@ -938,7 +930,7 @@ def test_insert_thousands_raises_exception_with_extra_columns_after_first_100(fr # https://github.com/simonw/sqlite-utils/issues/139 with pytest.raises(Exception, match="table test has no column named extra"): fresh_db["test"].insert_all( - [{"i": i, "word": "word_{}".format(i)} for i in range(100)] + [{"i": i, "word": f"word_{i}"} for i in range(100)] + [{"i": 101, "extra": "This extra column should cause an exception"}], ) @@ -946,7 +938,7 @@ def test_insert_thousands_raises_exception_with_extra_columns_after_first_100(fr def test_insert_thousands_adds_extra_columns_after_first_100_with_alter(fresh_db): # https://github.com/simonw/sqlite-utils/issues/139 fresh_db["test"].insert_all( - [{"i": i, "word": "word_{}".format(i)} for i in range(100)] + [{"i": i, "word": f"word_{i}"} for i in range(100)] + [{"i": 101, "extra": "Should trigger ALTER"}], alter=True, ) @@ -958,7 +950,7 @@ def test_insert_thousands_adds_extra_columns_after_first_100_with_alter(fresh_db def test_insert_all_pk_not_in_records_raises(fresh_db, num_rows): # https://github.com/simonw/sqlite-utils/issues/732 fresh_db.conn.execute("CREATE TABLE t (a TEXT, b INT, PRIMARY KEY (a, b))") - rows = [{"a": "x{}".format(i), "b": i} for i in range(num_rows)] + rows = [{"a": f"x{i}", "b": i} for i in range(num_rows)] with pytest.raises(InvalidColumns) as ex: fresh_db["t"].insert_all(rows, pk="not_a_column") @@ -975,7 +967,7 @@ def test_insert_all_pk_not_in_records_alter_raises(fresh_db, num_rows): # known - a pk column that is in neither the table nor the records # still raises fresh_db.conn.execute("CREATE TABLE t (a TEXT, b INT, PRIMARY KEY (a, b))") - rows = [{"a": "x{}".format(i), "b": i} for i in range(num_rows)] + rows = [{"a": f"x{i}", "b": i} for i in range(num_rows)] with pytest.raises(InvalidColumns) as ex: fresh_db["t"].insert_all(rows, pk="not_a_column", alter=True) @@ -1146,7 +1138,7 @@ def test_insert_hash_id_columns(fresh_db, use_table_factory): insert_kwargs = {} else: dogs = fresh_db["dogs"] - insert_kwargs = dict(hash_id_columns=("name", "twitter")) + insert_kwargs = {"hash_id_columns": ("name", "twitter")} id = dogs.insert( {"name": "Cleo", "twitter": "cleopaws", "age": 5}, @@ -1654,7 +1646,7 @@ def test_upsert_uses_pk_from_prior_insert_655(fresh_db): # Upsert should work without specifying pk again table.upsert({"id": 1, "name": "Alice Updated"}) assert table.count == 1 - assert list(table.rows)[0]["name"] == "Alice Updated" + assert next(iter(table.rows))["name"] == "Alice Updated" def test_upsert_all_uses_pk_from_prior_insert_655(fresh_db): diff --git a/tests/test_create_view.py b/tests/test_create_view.py index 056e246..2b70099 100644 --- a/tests/test_create_view.py +++ b/tests/test_create_view.py @@ -1,4 +1,5 @@ import pytest + from sqlite_utils.utils import OperationalError diff --git a/tests/test_default_value.py b/tests/test_default_value.py index 3724d99..2815180 100644 --- a/tests/test_default_value.py +++ b/tests/test_default_value.py @@ -31,7 +31,7 @@ EXAMPLES = [ @pytest.mark.parametrize("column_def,initial_value,expected_value", EXAMPLES) def test_quote_default_value(fresh_db, column_def, initial_value, expected_value): - fresh_db.execute("create table foo (col {})".format(column_def)) + fresh_db.execute(f"create table foo (col {column_def})") assert initial_value == fresh_db["foo"].columns[0].default_value assert expected_value == fresh_db.quote_default_value( fresh_db["foo"].columns[0].default_value diff --git a/tests/test_delete.py b/tests/test_delete.py index a2d93aa..dffb6bb 100644 --- a/tests/test_delete.py +++ b/tests/test_delete.py @@ -3,7 +3,7 @@ import sqlite_utils def test_delete_rowid_table(fresh_db): table = fresh_db["table"] - table.insert({"foo": 1}).last_pk + table.insert({"foo": 1}) rowid = table.insert({"foo": 2}).last_pk table.delete(rowid) assert [{"foo": 1}] == list(table.rows) diff --git a/tests/test_docs.py b/tests/test_docs.py index f657416..6bc06c8 100644 --- a/tests/test_docs.py +++ b/tests/test_docs.py @@ -1,8 +1,10 @@ -from click.testing import CliRunner -from sqlite_utils import cli, recipes -from pathlib import Path -import pytest import re +from pathlib import Path + +import pytest +from click.testing import CliRunner + +from sqlite_utils import cli, recipes docs_path = Path(__file__).parent.parent / "docs" commands_re = re.compile(r"(?:\$ | )sqlite-utils (\S+)") @@ -34,7 +36,7 @@ def test_commands_are_documented(documented_commands, command): @pytest.mark.parametrize("command", cli.cli.commands.values()) def test_commands_have_help(command): - assert command.help, "{} is missing its help".format(command) + assert command.help, f"{command} is missing its help" def test_convert_help(): diff --git a/tests/test_duplicate.py b/tests/test_duplicate.py index 28961d2..ad853a5 100644 --- a/tests/test_duplicate.py +++ b/tests/test_duplicate.py @@ -1,7 +1,9 @@ -from sqlite_utils.db import NoTable import datetime + import pytest +from sqlite_utils.db import NoTable + def test_duplicate(fresh_db): # Create table using native Sqlite statement: @@ -12,7 +14,7 @@ def test_duplicate(fresh_db): "bool_col" INTEGER, "datetime_col" TEXT)""") # Insert one row of mock data: - dt = datetime.datetime.now() + dt = datetime.datetime.now(datetime.timezone.utc) data = { "text_col": "Cleo", "real_col": 3.14, diff --git a/tests/test_enable_counts.py b/tests/test_enable_counts.py index 2f6b0db..71a8936 100644 --- a/tests/test_enable_counts.py +++ b/tests/test_enable_counts.py @@ -1,14 +1,14 @@ -from sqlite_utils import Database -from sqlite_utils import cli -from click.testing import CliRunner import pytest +from click.testing import CliRunner + +from sqlite_utils import Database, cli def test_enable_counts_specific_table(fresh_db): foo = fresh_db["foo"] assert fresh_db.table_names() == [] for i in range(10): - foo.insert({"name": "item {}".format(i)}) + foo.insert({"name": f"item {i}"}) assert fresh_db.table_names() == ["foo"] assert foo.count == 10 # Now enable counts @@ -44,7 +44,7 @@ def test_enable_counts_specific_table(fresh_db): assert list(fresh_db["_counts"].rows) == [{"count": 10, "table": "foo"}] # Add some items to test the triggers for i in range(5): - foo.insert({"name": "item {}".format(10 + i)}) + foo.insert({"name": f"item {10 + i}"}) assert foo.count == 15 assert list(fresh_db["_counts"].rows) == [{"count": 15, "table": "foo"}] # Delete some items diff --git a/tests/test_extract.py b/tests/test_extract.py index c73ee7a..915e6e1 100644 --- a/tests/test_extract.py +++ b/tests/test_extract.py @@ -1,19 +1,21 @@ -from sqlite_utils.db import InvalidColumns import itertools + import pytest +from sqlite_utils.db import InvalidColumns + @pytest.mark.parametrize("table", [None, "Species"]) @pytest.mark.parametrize("fk_column", [None, "species"]) def test_extract_single_column(fresh_db, table, fk_column): expected_table = table or "species" - expected_fk = fk_column or "{}_id".format(expected_table) + expected_fk = fk_column or f"{expected_table}_id" iter_species = itertools.cycle(["Palm", "Spruce", "Mangrove", "Oak"]) fresh_db["tree"].insert_all( ( { "id": i, - "name": "Tree {}".format(i), + "name": f"Tree {i}", "species": next(iter_species), "end": 1, } @@ -26,13 +28,12 @@ def test_extract_single_column(fresh_db, table, fk_column): 'CREATE TABLE "tree" (\n' ' "id" INTEGER PRIMARY KEY,\n' ' "name" TEXT,\n' - ' "{}" INTEGER REFERENCES "{}"("id"),\n'.format(expected_fk, expected_table) + f' "{expected_fk}" INTEGER REFERENCES "{expected_table}"("id"),\n' + ' "end" INTEGER\n' + ")" ) assert fresh_db[expected_table].schema == ( - 'CREATE TABLE "{}" (\n'.format(expected_table) - + ' "id" INTEGER PRIMARY KEY,\n' + f'CREATE TABLE "{expected_table}" (\n' + ' "id" INTEGER PRIMARY KEY,\n' ' "species" TEXT\n' ")" ) @@ -57,7 +58,7 @@ def test_extract_multiple_columns_with_rename(fresh_db): ( { "id": i, - "name": "Tree {}".format(i), + "name": f"Tree {i}", "common_name": next(iter_common), "latin_name": next(iter_latin), } diff --git a/tests/test_extracts.py b/tests/test_extracts.py index 7add79a..9519b91 100644 --- a/tests/test_extracts.py +++ b/tests/test_extracts.py @@ -1,13 +1,14 @@ -from sqlite_utils.db import Index import pytest +from sqlite_utils.db import Index + @pytest.mark.parametrize( "kwargs,expected_table", [ - (dict(extracts={"species_id": "Species"}), "Species"), - (dict(extracts=["species_id"]), "species_id"), - (dict(extracts=("species_id",)), "species_id"), + ({"extracts": {"species_id": "Species"}}, "Species"), + ({"extracts": ["species_id"]}, "species_id"), + ({"extracts": ("species_id",)}, "species_id"), ], ) @pytest.mark.parametrize("use_table_factory", [True, False]) @@ -30,15 +31,11 @@ def test_extracts(fresh_db, kwargs, expected_table, use_table_factory): # Should now have two tables: Trees and Species assert {expected_table, "Trees"} == set(fresh_db.table_names()) assert ( - 'CREATE TABLE "{}" (\n "id" INTEGER PRIMARY KEY,\n "value" TEXT\n)'.format( - expected_table - ) + f'CREATE TABLE "{expected_table}" (\n "id" INTEGER PRIMARY KEY,\n "value" TEXT\n)' == fresh_db[expected_table].schema ) assert ( - 'CREATE TABLE "Trees" (\n "id" INTEGER,\n "species_id" INTEGER REFERENCES "{}"("id")\n)'.format( - expected_table - ) + f'CREATE TABLE "Trees" (\n "id" INTEGER,\n "species_id" INTEGER REFERENCES "{expected_table}"("id")\n)' == fresh_db["Trees"].schema ) # Should have a foreign key reference @@ -51,7 +48,7 @@ def test_extracts(fresh_db, kwargs, expected_table, use_table_factory): assert [ Index( seq=0, - name="idx_{}_value".format(expected_table), + name=f"idx_{expected_table}_value", unique=1, origin="c", partial=0, diff --git a/tests/test_foreign_keys.py b/tests/test_foreign_keys.py index b37d374..45f4f35 100644 --- a/tests/test_foreign_keys.py +++ b/tests/test_foreign_keys.py @@ -1,6 +1,7 @@ """Tests for compound (multi-column) foreign keys - issue #594.""" import pytest + from sqlite_utils import Database from sqlite_utils.db import AlterError, ForeignKey from sqlite_utils.utils import sqlite3 @@ -64,7 +65,7 @@ def test_foreign_key_no_longer_unpacks_as_tuple(fresh_db): fresh_db["books"].add_foreign_key("author_id", "authors", "id") fk = fresh_db["books"].foreign_keys[0] with pytest.raises(TypeError): - table, column, other_table, other_column = fk + _table, _column, _other_table, _other_column = fk with pytest.raises(TypeError): fk[0] diff --git a/tests/test_fts.py b/tests/test_fts.py index 64ec645..50c1770 100644 --- a/tests/test_fts.py +++ b/tests/test_fts.py @@ -1,7 +1,9 @@ +from unittest.mock import ANY + import pytest + from sqlite_utils import Database from sqlite_utils.utils import sqlite3 -from unittest.mock import ANY search_records = [ { @@ -103,9 +105,10 @@ def test_search_limit_offset(fresh_db): table.enable_fts(["text", "country"], fts_version="FTS4") assert len(list(table.search("are"))) == 2 assert len(list(table.search("are", limit=1))) == 1 - assert list(table.search("are", limit=1, order_by="rowid"))[0]["rowid"] == 1 + assert next(iter(table.search("are", limit=1, order_by="rowid")))["rowid"] == 1 assert ( - list(table.search("are", limit=1, offset=1, order_by="rowid"))[0]["rowid"] == 2 + next(iter(table.search("are", limit=1, offset=1, order_by="rowid")))["rowid"] + == 2 ) @@ -223,20 +226,20 @@ def test_populate_fts_escape_table_names(fresh_db): @pytest.mark.parametrize("fts_version", ("4", "5")) def test_fts_tokenize(fresh_db, fts_version): - table_name = "searchable_{}".format(fts_version) + table_name = f"searchable_{fts_version}" table = fresh_db[table_name] table.insert_all(search_records) # Test without porter stemming table.enable_fts( ["text", "country"], - fts_version="FTS{}".format(fts_version), + fts_version=f"FTS{fts_version}", ) assert [] == list(table.search("bite")) # Test WITH stemming table.disable_fts() table.enable_fts( ["text", "country"], - fts_version="FTS{}".format(fts_version), + fts_version=f"FTS{fts_version}", tokenize="porter", ) rows = list(table.search("bite", order_by="rowid")) @@ -251,10 +254,10 @@ def test_fts_tokenize(fresh_db, fts_version): def test_optimize_fts(fresh_db): for fts_version in ("4", "5"): - table_name = "searchable_{}".format(fts_version) + table_name = f"searchable_{fts_version}" table = fresh_db[table_name] table.insert_all(search_records) - table.enable_fts(["text", "country"], fts_version="FTS{}".format(fts_version)) + table.enable_fts(["text", "country"], fts_version=f"FTS{fts_version}") # You can call optimize successfully against the tables OR their _fts equivalents: for table_name in ( "searchable_4", @@ -310,12 +313,12 @@ def test_disable_fts(fresh_db, create_triggers): expected_triggers = {"searchable_ai", "searchable_ad", "searchable_au"} else: expected_triggers = set() - assert expected_triggers == set( + assert expected_triggers == { r[0] for r in fresh_db.execute( "select name from sqlite_master where type = 'trigger'" ).fetchall() - ) + } # Now run .disable_fts() and confirm it worked table.disable_fts() assert ( @@ -424,7 +427,7 @@ def test_enable_fts_replace(kwargs): db["books"].enable_fts(**kwargs, replace=True) # Check that the new configuration is correct if should_have_changed_columns: - assert db["books_fts"].columns_dict.keys() == set(["title"]) + assert db["books_fts"].columns_dict.keys() == {"title"} if "create_triggers" in kwargs: assert db["books"].triggers if "fts_version" in kwargs: @@ -741,6 +744,7 @@ def test_enable_fts_cli_on_view_errors(tmpdir): db.create_view("v", "select * from t") db.close() from click.testing import CliRunner + from sqlite_utils import cli as cli_module result = CliRunner().invoke(cli_module.cli, ["enable-fts", db_path, "v", "text"]) diff --git a/tests/test_get.py b/tests/test_get.py index 63c4a2e..3cdaed8 100644 --- a/tests/test_get.py +++ b/tests/test_get.py @@ -1,4 +1,5 @@ import pytest + from sqlite_utils.db import NotFoundError diff --git a/tests/test_gis.py b/tests/test_gis.py index f39554e..8b41d22 100644 --- a/tests/test_gis.py +++ b/tests/test_gis.py @@ -1,7 +1,8 @@ import json -import pytest +import pytest from click.testing import CliRunner + from sqlite_utils.cli import cli from sqlite_utils.db import Database from sqlite_utils.utils import find_spatialite, sqlite3 @@ -104,7 +105,7 @@ def test_query_load_extension(use_spatialite_shortcut): [ ":memory:", "select spatialite_version()", - "--load-extension={}".format(load_extension), + f"--load-extension={load_extension}", ], ) assert result.exit_code == 0, result.stdout diff --git a/tests/test_hypothesis.py b/tests/test_hypothesis.py index f12f865..ab652c7 100644 --- a/tests/test_hypothesis.py +++ b/tests/test_hypothesis.py @@ -1,5 +1,6 @@ -from hypothesis import given import hypothesis.strategies as st +from hypothesis import given + import sqlite_utils diff --git a/tests/test_insert_files.py b/tests/test_insert_files.py index 88e49a8..1724d2d 100644 --- a/tests/test_insert_files.py +++ b/tests/test_insert_files.py @@ -1,10 +1,12 @@ -from sqlite_utils import cli, Database -from click.testing import CliRunner import os import pathlib -import pytest import sys +import pytest +from click.testing import CliRunner + +from sqlite_utils import Database, cli + @pytest.mark.parametrize("silent", (False, True)) @pytest.mark.parametrize( @@ -44,7 +46,7 @@ def test_insert_files(silent, pk_args, expected_pks): ) cols = [] for coltype in coltypes: - cols += ["-c", "{}:{}".format(coltype, coltype)] + cols += ["-c", f"{coltype}:{coltype}"] result = runner.invoke( cli.cli, ["insert-files", db_path, "files", str(tmpdir)] @@ -142,7 +144,7 @@ def test_insert_files_stdin(use_text, encoding, input, expected): ) assert result.exit_code == 0, result.stdout db = Database(db_path) - row = list(db["files"].rows)[0] + row = next(iter(db["files"].rows)) key = "content" if use_text: key = "content_text" @@ -167,5 +169,5 @@ def test_insert_files_bad_text_encoding_error(): ) assert result.exit_code == 1, result.output assert result.output.strip().startswith( - "Error: Could not read file '{}' as text".format(str(latin.resolve())) + f"Error: Could not read file '{latin.resolve()!s}' as text" ) diff --git a/tests/test_introspect.py b/tests/test_introspect.py index 8b6765d..385c052 100644 --- a/tests/test_introspect.py +++ b/tests/test_introspect.py @@ -1,6 +1,7 @@ -from sqlite_utils.db import Index, View, Database, XIndex, XIndexColumn import pytest +from sqlite_utils.db import Database, Index, View, XIndex, XIndexColumn + def _check_supports_strict(): """Check if SQLite supports strict tables without leaking the database.""" @@ -57,8 +58,8 @@ def test_detect_fts_similar_tables(fresh_db, reverse_order): fresh_db[table2].insert({"title": "Hello"}).enable_fts( ["title"], fts_version="FTS4" ) - assert fresh_db[table1].detect_fts() == "{}_fts".format(table1) - assert fresh_db[table2].detect_fts() == "{}_fts".format(table2) + assert fresh_db[table1].detect_fts() == f"{table1}_fts" + assert fresh_db[table2].detect_fts() == f"{table2}_fts" def test_tables(existing_db): diff --git a/tests/test_list_mode.py b/tests/test_list_mode.py index 746c9c1..646098e 100644 --- a/tests/test_list_mode.py +++ b/tests/test_list_mode.py @@ -3,6 +3,7 @@ Tests for list-based iteration in insert_all and upsert_all """ import pytest + from sqlite_utils import Database diff --git a/tests/test_lookup.py b/tests/test_lookup.py index da4f18b..c93d1ed 100644 --- a/tests/test_lookup.py +++ b/tests/test_lookup.py @@ -1,6 +1,7 @@ -from sqlite_utils.db import Index import pytest +from sqlite_utils.db import Index + def test_lookup_new_table(fresh_db): species = fresh_db["species"] diff --git a/tests/test_m2m.py b/tests/test_m2m.py index d613bb9..4fca918 100644 --- a/tests/test_m2m.py +++ b/tests/test_m2m.py @@ -1,6 +1,7 @@ -from sqlite_utils.db import ForeignKey, NoObviousTable import pytest +from sqlite_utils.db import ForeignKey, NoObviousTable + def test_insert_m2m_single(fresh_db): dogs = fresh_db["dogs"] @@ -65,8 +66,7 @@ def test_insert_m2m_iterable(fresh_db): iterable_records = ({"id": 1, "name": "Phineas"}, {"id": 2, "name": "Ferb"}) def iterable(): - for record in iterable_records: - yield record + yield from iterable_records platypuses = fresh_db["platypuses"] platypuses.insert({"id": 1, "name": "Perry"}, pk="id").m2m( diff --git a/tests/test_migrations.py b/tests/test_migrations.py index 04185fc..3f3dfea 100644 --- a/tests/test_migrations.py +++ b/tests/test_migrations.py @@ -1,4 +1,5 @@ import pytest + import sqlite_utils from sqlite_utils import Migrations @@ -154,10 +155,9 @@ def test_non_transactional_migration_allows_vacuum(tmpdir): def test_apply_composes_inside_outer_transaction(migrations): db = sqlite_utils.Database(memory=True) - with pytest.raises(ZeroDivisionError): - with db.atomic(): - migrations.apply(db) - raise ZeroDivisionError + with pytest.raises(ZeroDivisionError), db.atomic(): + migrations.apply(db) + raise ZeroDivisionError # The outer transaction rolled back, taking the migrations with it assert db.table_names() == [] diff --git a/tests/test_plugins.py b/tests/test_plugins.py index c793e32..ef202be 100644 --- a/tests/test_plugins.py +++ b/tests/test_plugins.py @@ -1,9 +1,12 @@ -from click.testing import CliRunner -import click import importlib -import pytest +import sqlite3 import sys -from sqlite_utils import cli, Database, hookimpl, plugins + +import click +import pytest +from click.testing import CliRunner + +from sqlite_utils import Database, cli, hookimpl, plugins def _supports_pragma_function_list(): @@ -11,7 +14,7 @@ def _supports_pragma_function_list(): try: db.execute("select * from pragma_function_list()") return True - except Exception: + except sqlite3.DatabaseError: return False finally: db.close() diff --git a/tests/test_query.py b/tests/test_query.py index 06847da..9d79755 100644 --- a/tests/test_query.py +++ b/tests/test_query.py @@ -1,6 +1,7 @@ -import pytest import types +import pytest + from sqlite_utils.utils import sqlite3 diff --git a/tests/test_recipes.py b/tests/test_recipes.py index a7c7ef7..c6222a3 100644 --- a/tests/test_recipes.py +++ b/tests/test_recipes.py @@ -1,7 +1,9 @@ +import json + +import pytest + from sqlite_utils import recipes from sqlite_utils.utils import sqlite3 -import json -import pytest @pytest.fixture diff --git a/tests/test_recreate.py b/tests/test_recreate.py index bce53d5..09e237e 100644 --- a/tests/test_recreate.py +++ b/tests/test_recreate.py @@ -1,8 +1,10 @@ -from sqlite_utils import Database -import sqlite3 import pathlib +import sqlite3 + import pytest +from sqlite_utils import Database + def test_recreate_ignored_for_in_memory(): # None of these should raise an exception: diff --git a/tests/test_rows_from_file.py b/tests/test_rows_from_file.py index a19fed6..8c080d6 100644 --- a/tests/test_rows_from_file.py +++ b/tests/test_rows_from_file.py @@ -1,7 +1,9 @@ -from sqlite_utils.utils import rows_from_file, Format, RowError from io import BytesIO, StringIO + import pytest +from sqlite_utils.utils import Format, RowError, rows_from_file + @pytest.mark.parametrize( "input,expected_format", @@ -29,7 +31,7 @@ def test_rows_from_file_detect_format(input, expected_format): ) def test_rows_from_file_extra_fields_strategies(ignore_extras, extras_key, expected): try: - rows, format = rows_from_file( + rows, _format = rows_from_file( BytesIO(b"id,name\r\n1,Cleo,oops"), format=Format.CSV, ignore_extras=ignore_extras, diff --git a/tests/test_sniff.py b/tests/test_sniff.py index 4bbdb66..7149978 100644 --- a/tests/test_sniff.py +++ b/tests/test_sniff.py @@ -1,7 +1,9 @@ -from sqlite_utils import cli, Database -from click.testing import CliRunner import pathlib + import pytest +from click.testing import CliRunner + +from sqlite_utils import Database, cli sniff_dir = pathlib.Path(__file__).parent / "sniff" diff --git a/tests/test_suggest_column_types.py b/tests/test_suggest_column_types.py index e36c58f..d4f28d3 100644 --- a/tests/test_suggest_column_types.py +++ b/tests/test_suggest_column_types.py @@ -1,5 +1,7 @@ -import pytest from collections import OrderedDict + +import pytest + from sqlite_utils.utils import suggest_column_types diff --git a/tests/test_tracer.py b/tests/test_tracer.py index d14697d..ec81f2f 100644 --- a/tests/test_tracer.py +++ b/tests/test_tracer.py @@ -53,16 +53,18 @@ def test_with_tracer(): assert len(collected) == 4 assert collected == [ ( - "SELECT name FROM sqlite_master\n" - " WHERE rootpage = 0\n" - " AND (\n" - " sql LIKE :like\n" - " OR sql LIKE :like2\n" - " OR (\n" - " tbl_name = :table\n" - " AND sql LIKE '%VIRTUAL TABLE%USING FTS%'\n" - " )\n" - " )", + ( + "SELECT name FROM sqlite_master\n" + " WHERE rootpage = 0\n" + " AND (\n" + " sql LIKE :like\n" + " OR sql LIKE :like2\n" + " OR (\n" + " tbl_name = :table\n" + " AND sql LIKE '%VIRTUAL TABLE%USING FTS%'\n" + " )\n" + " )" + ), { "like": "%VIRTUAL TABLE%USING FTS%content=[dogs]%", "like2": '%VIRTUAL TABLE%USING FTS%content="dogs"%', @@ -72,21 +74,23 @@ def test_with_tracer(): ("select name from sqlite_master where type = 'view'", None), ("select sql from sqlite_master where name = ?", ("dogs_fts",)), ( - 'with "original" as (\n' - " select\n" - " rowid,\n" - " *\n" - ' from "dogs"\n' - ")\n" - "select\n" - ' "original".*\n' - "from\n" - ' "original"\n' - ' join "dogs_fts" on "original".rowid = "dogs_fts".rowid\n' - "where\n" - ' "dogs_fts" match :query\n' - "order by\n" - ' "dogs_fts".rank', + ( + 'with "original" as (\n' + " select\n" + " rowid,\n" + " *\n" + ' from "dogs"\n' + ")\n" + "select\n" + ' "original".*\n' + "from\n" + ' "original"\n' + ' join "dogs_fts" on "original".rowid = "dogs_fts".rowid\n' + "where\n" + ' "dogs_fts" match :query\n' + "order by\n" + ' "dogs_fts".rank' + ), {"query": "Cleopaws"}, ), ] diff --git a/tests/test_transform.py b/tests/test_transform.py index 362f1ca..b9ee126 100644 --- a/tests/test_transform.py +++ b/tests/test_transform.py @@ -1,8 +1,9 @@ import sqlite3 +import pytest + from sqlite_utils.db import ForeignKey, TransactionError, TransformError from sqlite_utils.utils import OperationalError -import pytest @pytest.mark.parametrize( @@ -113,7 +114,7 @@ def test_transform_sql_table_with_primary_key( if use_pragma_foreign_keys: fresh_db.conn.execute("PRAGMA foreign_keys=ON") dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id") - sql = dogs.transform_sql(**{**params, **{"tmp_suffix": "suffix"}}) + sql = dogs.transform_sql(**{**params, "tmp_suffix": "suffix"}) assert sql == expected_sql # Check that .transform() runs without exceptions: with fresh_db.tracer(tracer): @@ -186,7 +187,7 @@ def test_transform_sql_table_with_no_primary_key( if use_pragma_foreign_keys: fresh_db.conn.execute("PRAGMA foreign_keys=ON") dogs.insert({"id": 1, "name": "Cleo", "age": "5"}) - sql = dogs.transform_sql(**{**params, **{"tmp_suffix": "suffix"}}) + sql = dogs.transform_sql(**{**params, "tmp_suffix": "suffix"}) assert sql == expected_sql # Check that .transform() runs without exceptions: with fresh_db.tracer(tracer): @@ -476,23 +477,22 @@ def test_transform_in_transaction_refuses_destructive_on_delete(fresh_db, on_del # keys inside an open transaction would fire those actions when the old # table is dropped - transform() should refuse instead fresh_db.conn.execute("PRAGMA foreign_keys=ON") - fresh_db.executescript(""" + fresh_db.executescript(f""" CREATE TABLE authors (id INTEGER PRIMARY KEY, name TEXT); CREATE TABLE books ( id INTEGER PRIMARY KEY, title TEXT, - author_id INTEGER REFERENCES authors(id) ON DELETE {} + author_id INTEGER REFERENCES authors(id) ON DELETE {on_delete} ); - """.format(on_delete)) + """) fresh_db["authors"].insert({"id": 1, "name": "Ursula K. Le Guin"}) fresh_db["books"].insert({"id": 1, "title": "The Dispossessed", "author_id": 1}) previous_schema = fresh_db["authors"].schema - with fresh_db.atomic(): - with pytest.raises(TransactionError) as excinfo: - fresh_db["authors"].transform(rename={"name": "author_name"}) + with fresh_db.atomic(), pytest.raises(TransactionError) as excinfo: + fresh_db["authors"].transform(rename={"name": "author_name"}) message = str(excinfo.value) assert "books" in message - assert "ON DELETE {}".format(on_delete.upper()) in message + assert f"ON DELETE {on_delete.upper()}" in message # Nothing should have changed assert fresh_db["authors"].schema == previous_schema assert list(fresh_db["books"].rows) == [ @@ -518,9 +518,8 @@ def test_transform_in_transaction_refuses_self_referential_cascade(fresh_db): {"id": 2, "name": "Science Fiction", "parent_id": 1}, ] ) - with fresh_db.atomic(): - with pytest.raises(TransactionError) as excinfo: - fresh_db["categories"].transform(rename={"name": "title"}) + with fresh_db.atomic(), pytest.raises(TransactionError) as excinfo: + fresh_db["categories"].transform(rename={"name": "title"}) assert "categories" in str(excinfo.value) assert fresh_db["categories"].count == 2 @@ -715,15 +714,15 @@ def test_transform_preserves_rowids(fresh_db, table_type): # Now delete and insert a row to mix up the `rowid` sequence fresh_db["places"].delete_where("id = ?", ["2"]) fresh_db["places"].insert({"id": "4", "name": "London", "country": "UK"}) - previous_rows = list( + previous_rows = [ tuple(row) for row in fresh_db.execute("select rowid, id, name from places") - ) + ] # Transform it fresh_db["places"].transform(column_order=("country", "name")) # Should be the same - next_rows = list( + next_rows = [ tuple(row) for row in fresh_db.execute("select rowid, id, name from places") - ) + ] assert previous_rows == next_rows diff --git a/tests/test_update.py b/tests/test_update.py index 03bec11..e6ae7d8 100644 --- a/tests/test_update.py +++ b/tests/test_update.py @@ -43,7 +43,7 @@ def test_update_compound_pk_table(fresh_db): ) def test_update_invalid_pk(fresh_db, pk, update_pk): table = fresh_db["table"] - table.insert({"id1": 5, "id2": 3, "v": 1}, pk=pk).last_pk + table.insert({"id1": 5, "id2": 3, "v": 1}, pk=pk) with pytest.raises(NotFoundError): table.update(update_pk, {"v": 2}) diff --git a/tests/test_upsert.py b/tests/test_upsert.py index a782b26..0eaae9b 100644 --- a/tests/test_upsert.py +++ b/tests/test_upsert.py @@ -1,7 +1,8 @@ -from sqlite_utils.db import PrimaryKeyRequired -from sqlite_utils import Database import pytest +from sqlite_utils import Database +from sqlite_utils.db import PrimaryKeyRequired + @pytest.mark.parametrize("use_old_upsert", (False, True)) def test_upsert(use_old_upsert): diff --git a/tests/test_utils.py b/tests/test_utils.py index 3de5e94..360a443 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -1,8 +1,10 @@ -from sqlite_utils import utils import csv import io + import pytest +from sqlite_utils import utils + @pytest.mark.parametrize( "input,expected,should_be_is", @@ -57,7 +59,7 @@ def test_maximize_csv_field_size_limit(): # Reset to default in case other tests have changed it csv.field_size_limit(utils.ORIGINAL_CSV_FIELD_SIZE_LIMIT) long_value = "a" * 131073 - long_csv = "id,text\n1,{}".format(long_value) + long_csv = f"id,text\n1,{long_value}" fp = io.BytesIO(long_csv.encode("utf-8")) # Using rows_from_file should error with pytest.raises(csv.Error): diff --git a/tests/test_wal.py b/tests/test_wal.py index 2ddcf54..35318f8 100644 --- a/tests/test_wal.py +++ b/tests/test_wal.py @@ -1,4 +1,5 @@ import pytest + from sqlite_utils import Database from sqlite_utils.db import TransactionError @@ -11,7 +12,7 @@ def db_path_tmpdir(tmpdir): def test_enable_disable_wal(db_path_tmpdir): - db, path, tmpdir = db_path_tmpdir + db, _path, tmpdir = db_path_tmpdir assert len(tmpdir.listdir()) == 1 assert "delete" == db.journal_mode assert "test.db-wal" not in [f.basename for f in tmpdir.listdir()] @@ -25,12 +26,11 @@ def test_enable_disable_wal(db_path_tmpdir): def test_enable_wal_inside_transaction_raises(db_path_tmpdir): - db, path, tmpdir = db_path_tmpdir + db, _path, _tmpdir = db_path_tmpdir db["test"].insert({"id": 1}, pk="id") - with pytest.raises(TransactionError): - with db.atomic(): - db["test"].insert({"id": 2}, pk="id") - db.enable_wal() + with pytest.raises(TransactionError), db.atomic(): + db["test"].insert({"id": 2}, pk="id") + db.enable_wal() # The atomic() block must have rolled back cleanly and the # journal mode must be unchanged assert db.journal_mode == "delete" @@ -38,19 +38,18 @@ def test_enable_wal_inside_transaction_raises(db_path_tmpdir): def test_disable_wal_inside_transaction_raises(db_path_tmpdir): - db, path, tmpdir = db_path_tmpdir + db, _path, _tmpdir = db_path_tmpdir db.enable_wal() db["test"].insert({"id": 1}, pk="id") - with pytest.raises(TransactionError): - with db.atomic(): - db["test"].insert({"id": 2}, pk="id") - db.disable_wal() + with pytest.raises(TransactionError), db.atomic(): + db["test"].insert({"id": 2}, pk="id") + db.disable_wal() assert db.journal_mode == "wal" assert [r["id"] for r in db["test"].rows] == [1] def test_ensure_autocommit_on(db_path_tmpdir): - db, path, tmpdir = db_path_tmpdir + db, _path, _tmpdir = db_path_tmpdir previous_isolation_level = db.conn.isolation_level assert previous_isolation_level is not None with db.ensure_autocommit_on(): @@ -63,7 +62,7 @@ def test_ensure_autocommit_on(db_path_tmpdir): def test_enable_wal_noop_inside_transaction_is_allowed(db_path_tmpdir): # Calling enable_wal() when WAL is already enabled is a no-op, # so it is fine inside a transaction - db, path, tmpdir = db_path_tmpdir + db, _path, _tmpdir = db_path_tmpdir db.enable_wal() with db.atomic(): db["test"].insert({"id": 1}, pk="id") @@ -75,13 +74,12 @@ def test_ensure_autocommit_on_inside_transaction_raises(db_path_tmpdir): # Setting isolation_level commits any pending transaction as a side # effect, silently breaking the caller's rollback guarantee - so # entering autocommit mode with a transaction open is an error - db, path, tmpdir = db_path_tmpdir + db, _path, _tmpdir = db_path_tmpdir db["test"].insert({"id": 1}, pk="id") db.begin() db.execute("insert into test (id) values (2)") - with pytest.raises(TransactionError): - with db.ensure_autocommit_on(): - pass + with pytest.raises(TransactionError), db.ensure_autocommit_on(): + pass # The transaction is still open and can still be rolled back assert db.conn.in_transaction db.rollback() From c621499ed1e3572989087c19a2f9d13bfff45021 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sat, 25 Jul 2026 14:53:46 -0700 Subject: [PATCH 6/6] codespell should check sqlite_utils as well It did in CI but did not in the Justfile --- Justfile | 1 + 1 file changed, 1 insertion(+) diff --git a/Justfile b/Justfile index 5caa120..be41523 100644 --- a/Justfile +++ b/Justfile @@ -16,6 +16,7 @@ uv run ty check sqlite_utils uv run cog --check README.md docs/*.rst uv run --group docs codespell docs/*.rst --ignore-words docs/codespell-ignore-words.txt + uv run --group docs codespell sqlite_utils --ignore-words docs/codespell-ignore-words.txt # Rebuild docs with cog @cog: