From cf2cb244faf992118f34aa196387a4ef8b39a20f Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Mon, 7 Sep 2020 14:56:59 -0700 Subject: [PATCH] Tracer mechanism for showing underlying SQL queries * Pass a tracer= function to Database constructor * New db.tracer() contextmanager * Neater SQL indentation, because tracer means it could be visible now * New db.execute() and db.executescript() methods Closes #150 --- docs/python-api.rst | 53 ++++++++++- sqlite_utils/cli.py | 2 +- sqlite_utils/db.py | 163 +++++++++++++++++++++------------- tests/conftest.py | 2 +- tests/test_cli.py | 4 +- tests/test_column_affinity.py | 2 +- tests/test_constructor.py | 4 +- tests/test_create.py | 2 +- tests/test_create_view.py | 10 +-- tests/test_fts.py | 4 +- tests/test_introspect.py | 2 +- tests/test_tracer.py | 62 +++++++++++++ 12 files changed, 231 insertions(+), 79 deletions(-) create mode 100644 tests/test_tracer.py diff --git a/docs/python-api.rst b/docs/python-api.rst index 134c824..5d3409d 100644 --- a/docs/python-api.rst +++ b/docs/python-api.rst @@ -43,6 +43,53 @@ Connections use ``PRAGMA recursive_triggers=on`` by default. If you don't want t db = Database(memory=True, recursive_triggers=False) +.. _python_api_tracing: + +Tracing queries +--------------- + +You can use the ``tracer`` mechanism to see SQL queries that are being executed by SQLite. A tracer is a function that you provide which will be called with ``sql`` and ``params`` arguments every time SQL is executed, for example: + +.. code-block:: python + + def tracer(sql, params): + print("SQL: {} - params: {}".format(sql, params)) + +You can pass this function to the ``Database()`` constructor like so: + +.. code-block:: python + + db = Database(memory=True, tracer=tracer) + +You can also turn on a tracer function temporarily for a block of code using the ``with db.tracer(...)`` context manager: + +.. code-block:: python + + db = Database(memory=True) + # ... later + with db.tracer(tracer): + db["dogs"].insert({"name": "Cleo"}) + +Queries will be passed to your ``tracer()`` function only for the duration of the ``with`` block. + +.. _python_api_execute: + +Executing queries +================= + +The ``db.execute()`` and ``db.executescript()`` methods provide wrappers around ``.execute()`` and ``.executescript()`` on the underlying SQLite connection. These wrappers log to the tracer function if one has been registered. + +.. code-block:: python + + db = Database(memory=True) + db["dogs"].insert({"name": "Cleo"}) + db.execute("update dogs set name = 'Cleopaws'") + +.. _python_api_table: + +Accessing tables +================ + Tables are accessed using the indexing operator, like so: .. code-block:: python @@ -930,7 +977,7 @@ For example: "postalCode": "95018" } }) - db.conn.execute(""" + db.execute(""" select json_extract(address, '$.addressLocality') from niche_museums """).fetchall() @@ -979,8 +1026,8 @@ A more useful example: if you are working with `SpatiaLite ".format(self.conn) + def execute(self, sql, parameters=None): + if self._tracer: + self._tracer(sql, parameters) + if parameters is not None: + return self.conn.execute(sql, parameters) + else: + return self.conn.execute(sql) + + def executescript(self, sql): + if self._tracer: + self._tracer(sql, None) + return self.conn.executescript(sql) + def table(self, table_name, **kwargs): klass = View if table_name in self.view_names() else Table return klass(self, table_name, **kwargs) @@ -139,7 +164,7 @@ class Database: # Normally we would use .execute(sql, [params]) for escaping, but # occasionally that isn't available - most notable when we need # to include a "... DEFAULT 'value'" in a column definition. - return self.conn.execute( + return self.execute( # Use SQLite itself to correctly escape this string: "SELECT quote(:value)", {"value": value}, @@ -152,12 +177,12 @@ class Database: if fts5: where.append("sql like '%FTS5%'") sql = "select name from sqlite_master where {}".format(" AND ".join(where)) - return [r[0] for r in self.conn.execute(sql).fetchall()] + return [r[0] for r in self.execute(sql).fetchall()] def view_names(self): return [ r[0] - for r in self.conn.execute( + for r in self.execute( "select name from sqlite_master where type = 'view'" ).fetchall() ] @@ -174,25 +199,25 @@ class Database: def triggers(self): return [ Trigger(*r) - for r in self.conn.execute( + for r in self.execute( "select name, tbl_name, sql from sqlite_master where type = 'trigger'" ).fetchall() ] @property def journal_mode(self): - return self.conn.execute("PRAGMA journal_mode;").fetchone()[0] + return self.execute("PRAGMA journal_mode;").fetchone()[0] def enable_wal(self): if self.journal_mode != "wal": - self.conn.execute("PRAGMA journal_mode=wal;") + self.execute("PRAGMA journal_mode=wal;") def disable_wal(self): if self.journal_mode != "delete": - self.conn.execute("PRAGMA journal_mode=delete;") + self.execute("PRAGMA journal_mode=delete;") def execute_returning_dicts(self, sql, params=None): - cursor = self.conn.execute(sql, params or tuple()) + cursor = self.execute(sql, params or tuple()) keys = [d[0] for d in cursor.description] return [dict(zip(keys, row)) for row in cursor.fetchall()] @@ -337,7 +362,7 @@ class Database: """.format( table=name, columns_sql=columns_sql, extra_pk=extra_pk ) - self.conn.execute(sql) + self.execute(sql) return self.table( name, pk=pk, @@ -363,7 +388,7 @@ class Database: if create_sql == self[name].schema: return self self[name].drop() - self.conn.execute(create_sql) + self.execute(create_sql) return self def m2m_table_candidates(self, table, other_table): @@ -451,7 +476,7 @@ class Database: table.create_index([fk.column]) def vacuum(self): - self.conn.execute("VACUUM;") + self.execute("VACUUM;") class Queryable: @@ -464,7 +489,7 @@ class Queryable: @property def count(self): - return self.db.conn.execute( + return self.db.execute( "select count(*) from [{}]".format(self.name) ).fetchone()[0] @@ -480,7 +505,7 @@ class Queryable: sql += " where " + where if order_by is not None: sql += " order by " + order_by - cursor = self.db.conn.execute(sql, where_args or []) + cursor = self.db.execute(sql, where_args or []) columns = [c[0] for c in cursor.description] for row in cursor: yield dict(zip(columns, row)) @@ -489,9 +514,7 @@ class Queryable: def columns(self): if not self.exists(): return [] - rows = self.db.conn.execute( - "PRAGMA table_info([{}])".format(self.name) - ).fetchall() + rows = self.db.execute("PRAGMA table_info([{}])".format(self.name)).fetchall() return [Column(*row) for row in rows] @property @@ -501,7 +524,7 @@ class Queryable: @property def schema(self): - return self.db.conn.execute( + return self.db.execute( "select sql from sqlite_master where name = ?", (self.name,) ).fetchone()[0] @@ -587,7 +610,7 @@ class Table(Queryable): @property def foreign_keys(self): fks = [] - for row in self.db.conn.execute( + for row in self.db.execute( "PRAGMA foreign_key_list([{}])".format(self.name) ).fetchall(): if row is not None: @@ -615,7 +638,7 @@ class Table(Queryable): ) column_sql = "PRAGMA index_info({})".format(index_name_quoted) columns = [] - for seqno, cid, name in self.db.conn.execute(column_sql).fetchall(): + for seqno, cid, name in self.db.execute(column_sql).fetchall(): columns.append(name) row["columns"] = columns # These columns may be missing on older SQLite versions: @@ -629,7 +652,7 @@ class Table(Queryable): def triggers(self): return [ Trigger(*r) - for r in self.db.conn.execute( + for r in self.db.execute( "select name, tbl_name, sql from sqlite_master where type = 'trigger'" " and tbl_name = ?", (self.name,), @@ -683,7 +706,7 @@ class Table(Queryable): if_not_exists="IF NOT EXISTS " if if_not_exists else "", ) ) - self.db.conn.execute(sql) + self.db.execute(sql) return self def add_column( @@ -720,13 +743,13 @@ class Table(Queryable): col_type=fk_col_type or COLUMN_TYPE_MAPPING[col_type], not_null_default=(" " + not_null_sql) if not_null_sql else "", ) - self.db.conn.execute(sql) + self.db.execute(sql) if fk is not None: self.add_foreign_key(col_name, fk, fk_col) return self def drop(self): - self.db.conn.execute("DROP TABLE [{}]".format(self.name)) + self.db.execute("DROP TABLE [{}]".format(self.name)) def guess_foreign_table(self, column): column = column.lower() @@ -811,7 +834,7 @@ class Table(Queryable): tokenize="\n tokenize='{}',".format(tokenize) if tokenize else "", ) ) - self.db.conn.executescript(sql) + self.db.executescript(sql) self.populate_fts(columns) if create_triggers: @@ -840,17 +863,23 @@ class Table(Queryable): new_cols=new_cols, ) ) - self.db.conn.executescript(triggers) + self.db.executescript(triggers) return self def populate_fts(self, columns): - sql = """ + sql = ( + textwrap.dedent( + """ INSERT INTO [{table}_fts] (rowid, {columns}) SELECT rowid, {columns} FROM [{table}]; - """.format( - table=self.name, columns=", ".join("[{}]".format(c) for c in columns) + """ + ) + .strip() + .format( + table=self.name, columns=", ".join("[{}]".format(c) for c in columns) + ) ) - self.db.conn.executescript(sql) + self.db.executescript(sql) return self def disable_fts(self): @@ -858,23 +887,29 @@ class Table(Queryable): if fts_table: self.db[fts_table].drop() # Now delete the triggers that related to that table - sql = """ + sql = ( + textwrap.dedent( + """ SELECT name FROM sqlite_master WHERE type = 'trigger' AND sql LIKE '% INSERT INTO [{}]%' - """.format( - fts_table + """ + ) + .strip() + .format(fts_table) ) trigger_names = [] - for row in self.db.conn.execute(sql).fetchall(): + for row in self.db.execute(sql).fetchall(): trigger_names.append(row[0]) with self.db.conn: for trigger_name in trigger_names: - self.db.conn.execute("DROP TRIGGER IF EXISTS [{}]".format(trigger_name)) + self.db.execute("DROP TRIGGER IF EXISTS [{}]".format(trigger_name)) def detect_fts(self): "Detect if table has a corresponding FTS virtual table and return it" - sql = """ + sql = ( + textwrap.dedent( + """ SELECT name FROM sqlite_master WHERE rootpage = 0 AND ( @@ -884,10 +919,12 @@ class Table(Queryable): AND sql LIKE '%VIRTUAL TABLE%USING FTS%' ) ) - """.format( - table=self.name + """ + ) + .strip() + .format(table=self.name) ) - rows = self.db.conn.execute(sql).fetchall() + rows = self.db.execute(sql).fetchall() if len(rows) == 0: return None else: @@ -896,18 +933,22 @@ class Table(Queryable): def optimize(self): fts_table = self.detect_fts() if fts_table is not None: - self.db.conn.execute( + self.db.execute( """ INSERT INTO [{table}] ([{table}]) VALUES ("optimize"); - """.format( + """.strip().format( table=fts_table ) ) - self.db.conn.execute( - """ + self.db.execute( + textwrap.dedent( + """ DELETE FROM [{table}_docsize] WHERE {column} NOT IN ( SELECT rowid FROM [{table}]); - """.format( + """ + ) + .strip() + .format( # FTS5 uses 'id' but FTS4 uses 'docid' column=self.db["{}_docsize".format(fts_table)].columns[0].name, table=fts_table, @@ -916,16 +957,20 @@ class Table(Queryable): return self def search(self, q): - sql = """ + sql = ( + textwrap.dedent( + """ select * from "{table}" where rowid in ( select rowid from [{table}_fts] where [{table}_fts] match :search ) order by rowid - """.format( - table=self.name + """ + ) + .strip() + .format(table=self.name) ) - return self.db.conn.execute(sql, (q,)).fetchall() + return self.db.execute(sql, (q,)).fetchall() def value_or_default(self, key, value): return self._defaults[key] if value is DEFAULT else value @@ -939,7 +984,7 @@ class Table(Queryable): table=self.name, wheres=" and ".join(wheres) ) with self.db.conn: - self.db.conn.execute(sql, pk_values) + self.db.execute(sql, pk_values) def delete_where(self, where=None, where_args=None): if not self.exists(): @@ -947,7 +992,7 @@ class Table(Queryable): sql = "delete from [{}]".format(self.name) if where is not None: sql += " where " + where - self.db.conn.execute(sql, where_args or []) + self.db.execute(sql, where_args or []) def update(self, pk_values, updates=None, alter=False, conversions=None): updates = updates or {} @@ -972,12 +1017,12 @@ class Table(Queryable): ) with self.db.conn: try: - rowcount = self.db.conn.execute(sql, args).rowcount + rowcount = self.db.execute(sql, args).rowcount except OperationalError as e: if alter and (" column" in e.args[0]): # Attempt to add any missing columns, then try again self.add_missing_columns([updates]) - rowcount = self.db.conn.execute(sql, args).rowcount + rowcount = self.db.execute(sql, args).rowcount else: raise @@ -1084,7 +1129,7 @@ class Table(Queryable): self.last_rowid = None self.last_pk = None if truncate and self.exists(): - self.db.conn.execute("DELETE FROM [{}];".format(self.name)) + self.db.execute("DELETE FROM [{}];".format(self.name)) for chunk in chunks(itertools.chain([first_record], records), batch_size): chunk = list(chunk) num_records_processed += len(chunk) @@ -1184,14 +1229,12 @@ class Table(Queryable): or_what = "OR IGNORE " sql = """ INSERT {or_what}INTO [{table}] ({columns}) VALUES {rows}; - """.format( + """.strip().format( or_what=or_what, table=self.name, columns=", ".join("[{}]".format(c) for c in all_columns), rows=", ".join( - """ - ({placeholders}) - """.format( + "({placeholders})".format( placeholders=", ".join( [conversions.get(col, "?") for col in all_columns] ) @@ -1205,12 +1248,12 @@ class Table(Queryable): with self.db.conn: for query, params in queries_and_params: try: - result = self.db.conn.execute(query, params) + result = self.db.execute(query, params) except OperationalError as e: if alter and (" column" in e.args[0]): # Attempt to add any missing columns, then try again self.add_missing_columns(chunk) - result = self.db.conn.execute(query, params) + result = self.db.execute(query, params) else: raise if num_records_processed == 1 and not upsert: @@ -1383,7 +1426,7 @@ class View(Queryable): ) def drop(self): - self.db.conn.execute("DROP VIEW [{}]".format(self.name)) + self.db.execute("DROP VIEW [{}]".format(self.name)) def chunks(sequence, size): diff --git a/tests/conftest.py b/tests/conftest.py index 1e1fc5e..4541689 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -10,7 +10,7 @@ def fresh_db(): @pytest.fixture def existing_db(): database = Database(memory=True) - database.conn.executescript( + database.executescript( """ CREATE TABLE foo (text TEXT); INSERT INTO foo (text) values ("one"); diff --git a/tests/test_cli.py b/tests/test_cli.py index 49631c9..964b6a2 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -396,7 +396,7 @@ def test_enable_fts_with_triggers(db_path): def search(q): return ( Database(db_path) - .conn.execute("select c1 from Gosh_fts where c1 match ?", [q]) + .execute("select c1 from Gosh_fts where c1 match ?", [q]) .fetchall() ) @@ -417,7 +417,7 @@ def test_populate_fts(db_path): def search(q): return ( Database(db_path) - .conn.execute("select c1 from Gosh_fts where c1 match ?", [q]) + .execute("select c1 from Gosh_fts where c1 match ?", [q]) .fetchall() ) diff --git a/tests/test_column_affinity.py b/tests/test_column_affinity.py index 4a34b25..fb8f340 100644 --- a/tests/test_column_affinity.py +++ b/tests/test_column_affinity.py @@ -41,5 +41,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.conn.execute("create table foo (col {})".format(column_def)) + fresh_db.execute("create table foo (col {})".format(column_def)) assert {"col": expected_type} == fresh_db["foo"].columns_dict diff --git a/tests/test_constructor.py b/tests/test_constructor.py index 8a790c2..b3cd963 100644 --- a/tests/test_constructor.py +++ b/tests/test_constructor.py @@ -4,9 +4,9 @@ import pytest def test_recursive_triggers(): db = Database(memory=True) - assert db.conn.execute("PRAGMA recursive_triggers").fetchone()[0] + assert db.execute("PRAGMA recursive_triggers").fetchone()[0] def test_recursive_triggers_off(): db = Database(memory=True, recursive_triggers=False) - assert not db.conn.execute("PRAGMA recursive_triggers").fetchone()[0] + assert not db.execute("PRAGMA recursive_triggers").fetchone()[0] diff --git a/tests/test_create.py b/tests/test_create.py index f6fad00..30a3d0e 100644 --- a/tests/test_create.py +++ b/tests/test_create.py @@ -680,7 +680,7 @@ def test_create_index_if_not_exists(fresh_db): ) def test_insert_dictionaries_and_lists_as_json(fresh_db, data_structure): fresh_db["test"].insert({"id": 1, "data": data_structure}, pk="id") - row = fresh_db.conn.execute("select id, data from test").fetchone() + row = fresh_db.execute("select id, data from test").fetchone() assert row[0] == 1 assert data_structure == json.loads(row[1]) diff --git a/tests/test_create_view.py b/tests/test_create_view.py index 26cf420..f0e2855 100644 --- a/tests/test_create_view.py +++ b/tests/test_create_view.py @@ -4,7 +4,7 @@ from sqlite_utils.utils import OperationalError def test_create_view(fresh_db): fresh_db.create_view("bar", "select 1 + 1") - rows = fresh_db.conn.execute("select * from bar").fetchall() + rows = fresh_db.execute("select * from bar").fetchall() assert [(2,)] == rows @@ -23,7 +23,7 @@ def test_create_view_ignore(fresh_db): fresh_db.create_view("bar", "select 1 + 1").create_view( "bar", "select 1 + 2", ignore=True ) - rows = fresh_db.conn.execute("select * from bar").fetchall() + rows = fresh_db.execute("select * from bar").fetchall() assert [(2,)] == rows @@ -31,13 +31,13 @@ def test_create_view_replace(fresh_db): fresh_db.create_view("bar", "select 1 + 1").create_view( "bar", "select 1 + 2", replace=True ) - rows = fresh_db.conn.execute("select * from bar").fetchall() + rows = fresh_db.execute("select * from bar").fetchall() assert [(3,)] == rows def test_create_view_replace_with_same_does_nothing(fresh_db): fresh_db.create_view("bar", "select 1 + 1") - initial_version = fresh_db.conn.execute("PRAGMA schema_version").fetchone()[0] + initial_version = fresh_db.execute("PRAGMA schema_version").fetchone()[0] fresh_db.create_view("bar", "select 1 + 1", replace=True) - after_version = fresh_db.conn.execute("PRAGMA schema_version").fetchone()[0] + after_version = fresh_db.execute("PRAGMA schema_version").fetchone()[0] assert after_version == initial_version diff --git a/tests/test_fts.py b/tests/test_fts.py index b1a946d..defb886 100644 --- a/tests/test_fts.py +++ b/tests/test_fts.py @@ -160,7 +160,7 @@ def test_disable_fts(fresh_db, create_triggers): expected_triggers = set() assert expected_triggers == set( r[0] - for r in fresh_db.conn.execute( + for r in fresh_db.execute( "select name from sqlite_master where type = 'trigger'" ).fetchall() ) @@ -168,7 +168,7 @@ def test_disable_fts(fresh_db, create_triggers): table.disable_fts() assert ( 0 - == fresh_db.conn.execute( + == fresh_db.execute( "select count(*) from sqlite_master where type = 'trigger'" ).fetchone()[0] ) diff --git a/tests/test_introspect.py b/tests/test_introspect.py index bebe537..cbec0a1 100644 --- a/tests/test_introspect.py +++ b/tests/test_introspect.py @@ -73,7 +73,7 @@ def test_table_repr(fresh_db): def test_indexes(fresh_db): - fresh_db.conn.executescript( + fresh_db.executescript( """ create table Gosh (c1 text, c2 text, c3 text); create index Gosh_c1 on Gosh(c1); diff --git a/tests/test_tracer.py b/tests/test_tracer.py new file mode 100644 index 0000000..e23ae86 --- /dev/null +++ b/tests/test_tracer.py @@ -0,0 +1,62 @@ +import pytest +from sqlite_utils import Database + + +def test_tracer(): + collected = [] + db = Database( + memory=True, tracer=lambda sql, params: collected.append((sql, params)) + ) + db["dogs"].insert({"name": "Cleopaws"}) + db["dogs"].enable_fts(["name"]) + db["dogs"].search("Cleopaws") + assert collected == [ + ("PRAGMA recursive_triggers=on;", None), + ("select name from sqlite_master where type = 'view'", None), + ("select name from sqlite_master where type = 'table'", None), + ("CREATE TABLE [dogs] (\n [name] TEXT\n);\n ", None), + ("select name from sqlite_master where type = 'view'", None), + ("INSERT INTO [dogs] ([name]) VALUES (?);", ["Cleopaws"]), + ("select name from sqlite_master where type = 'view'", None), + ( + "CREATE VIRTUAL TABLE [dogs_fts] USING FTS5 (\n [name],\n content=[dogs]\n);", + None, + ), + ( + "INSERT INTO [dogs_fts] (rowid, [name])\n SELECT rowid, [name] FROM [dogs];", + None, + ), + ("select name from sqlite_master where type = 'view'", None), + ( + 'select * from "dogs" where rowid in (\n select rowid from [dogs_fts]\n where [dogs_fts] match :search\n)\norder by rowid', + ("Cleopaws",), + ), + ] + + +def test_with_tracer(): + collected = [] + tracer = lambda sql, params: collected.append((sql, params)) + + db = Database(memory=True) + + db["dogs"].insert({"name": "Cleopaws"}) + db["dogs"].enable_fts(["name"]) + + assert len(collected) == 0 + + with db.tracer(tracer): + db["dogs"].search("Cleopaws") + + assert len(collected) == 2 + assert collected == [ + ("select name from sqlite_master where type = 'view'", None), + ( + 'select * from "dogs" where rowid in (\n select rowid from [dogs_fts]\n where [dogs_fts] match :search\n)\norder by rowid', + ("Cleopaws",), + ), + ] + + # Outside the with block collected should not be appended to + db["dogs"].insert({"name": "Cleopaws"}) + assert len(collected) == 2