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
This commit is contained in:
Simon Willison 2020-09-07 14:56:59 -07:00 committed by GitHub
commit cf2cb244fa
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
12 changed files with 231 additions and 79 deletions

View file

@ -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) 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: Tables are accessed using the indexing operator, like so:
.. code-block:: python .. code-block:: python
@ -930,7 +977,7 @@ For example:
"postalCode": "95018" "postalCode": "95018"
} }
}) })
db.conn.execute(""" db.execute("""
select json_extract(address, '$.addressLocality') select json_extract(address, '$.addressLocality')
from niche_museums from niche_museums
""").fetchall() """).fetchall()
@ -979,8 +1026,8 @@ A more useful example: if you are working with `SpatiaLite <https://www.gaia-gis
places = db["places"].create({"id": int, "name": str,}) places = db["places"].create({"id": int, "name": str,})
# Add a SpatiaLite 'geometry' column: # Add a SpatiaLite 'geometry' column:
db.conn.execute("select InitSpatialMetadata(1)") db.execute("select InitSpatialMetadata(1)")
db.conn.execute( db.execute(
"SELECT AddGeometryColumn('places', 'geometry', 4326, 'MULTIPOLYGON', 2);" "SELECT AddGeometryColumn('places', 'geometry', 4326, 'MULTIPOLYGON', 2);"
) )

View file

@ -771,7 +771,7 @@ def query(
for ext in load_extension: for ext in load_extension:
db.conn.load_extension(ext) db.conn.load_extension(ext)
with db.conn: with db.conn:
cursor = db.conn.execute(sql, dict(param)) cursor = db.execute(sql, dict(param))
if cursor.description is None: if cursor.description is None:
# This was an update/insert # This was an update/insert
headers = ["rows_affected"] headers = ["rows_affected"]

View file

@ -1,5 +1,6 @@
from .utils import sqlite3, OperationalError, suggest_column_types, column_affinity from .utils import sqlite3, OperationalError, suggest_column_types, column_affinity
from collections import namedtuple, OrderedDict from collections import namedtuple, OrderedDict
import contextlib
import datetime import datetime
import decimal import decimal
import hashlib import hashlib
@ -109,6 +110,7 @@ class Database:
memory=False, memory=False,
recreate=False, recreate=False,
recursive_triggers=True, recursive_triggers=True,
tracer=None,
): ):
assert (filename_or_conn is not None and not memory) or ( assert (filename_or_conn is not None and not memory) or (
filename_or_conn is None and memory filename_or_conn is None and memory
@ -122,8 +124,18 @@ class Database:
else: else:
assert not recreate, "recreate cannot be used with connections, only paths" assert not recreate, "recreate cannot be used with connections, only paths"
self.conn = filename_or_conn self.conn = filename_or_conn
self._tracer = tracer
if recursive_triggers: if recursive_triggers:
self.conn.execute("PRAGMA recursive_triggers=on;") self.execute("PRAGMA recursive_triggers=on;")
@contextlib.contextmanager
def tracer(self, tracer=None):
prev_tracer = self._tracer
self._tracer = tracer or print
try:
yield self
finally:
self._tracer = prev_tracer
def __getitem__(self, table_name): def __getitem__(self, table_name):
return self.table(table_name) return self.table(table_name)
@ -131,6 +143,19 @@ class Database:
def __repr__(self): def __repr__(self):
return "<Database {}>".format(self.conn) return "<Database {}>".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): def table(self, table_name, **kwargs):
klass = View if table_name in self.view_names() else Table klass = View if table_name in self.view_names() else Table
return klass(self, table_name, **kwargs) return klass(self, table_name, **kwargs)
@ -139,7 +164,7 @@ class Database:
# Normally we would use .execute(sql, [params]) for escaping, but # Normally we would use .execute(sql, [params]) for escaping, but
# occasionally that isn't available - most notable when we need # occasionally that isn't available - most notable when we need
# to include a "... DEFAULT 'value'" in a column definition. # to include a "... DEFAULT 'value'" in a column definition.
return self.conn.execute( return self.execute(
# Use SQLite itself to correctly escape this string: # Use SQLite itself to correctly escape this string:
"SELECT quote(:value)", "SELECT quote(:value)",
{"value": value}, {"value": value},
@ -152,12 +177,12 @@ class Database:
if fts5: if fts5:
where.append("sql like '%FTS5%'") where.append("sql like '%FTS5%'")
sql = "select name from sqlite_master where {}".format(" AND ".join(where)) 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): def view_names(self):
return [ return [
r[0] r[0]
for r in self.conn.execute( for r in self.execute(
"select name from sqlite_master where type = 'view'" "select name from sqlite_master where type = 'view'"
).fetchall() ).fetchall()
] ]
@ -174,25 +199,25 @@ class Database:
def triggers(self): def triggers(self):
return [ return [
Trigger(*r) Trigger(*r)
for r in self.conn.execute( for r in self.execute(
"select name, tbl_name, sql from sqlite_master where type = 'trigger'" "select name, tbl_name, sql from sqlite_master where type = 'trigger'"
).fetchall() ).fetchall()
] ]
@property @property
def journal_mode(self): 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): def enable_wal(self):
if self.journal_mode != "wal": if self.journal_mode != "wal":
self.conn.execute("PRAGMA journal_mode=wal;") self.execute("PRAGMA journal_mode=wal;")
def disable_wal(self): def disable_wal(self):
if self.journal_mode != "delete": 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): 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] keys = [d[0] for d in cursor.description]
return [dict(zip(keys, row)) for row in cursor.fetchall()] return [dict(zip(keys, row)) for row in cursor.fetchall()]
@ -337,7 +362,7 @@ class Database:
""".format( """.format(
table=name, columns_sql=columns_sql, extra_pk=extra_pk table=name, columns_sql=columns_sql, extra_pk=extra_pk
) )
self.conn.execute(sql) self.execute(sql)
return self.table( return self.table(
name, name,
pk=pk, pk=pk,
@ -363,7 +388,7 @@ class Database:
if create_sql == self[name].schema: if create_sql == self[name].schema:
return self return self
self[name].drop() self[name].drop()
self.conn.execute(create_sql) self.execute(create_sql)
return self return self
def m2m_table_candidates(self, table, other_table): def m2m_table_candidates(self, table, other_table):
@ -451,7 +476,7 @@ class Database:
table.create_index([fk.column]) table.create_index([fk.column])
def vacuum(self): def vacuum(self):
self.conn.execute("VACUUM;") self.execute("VACUUM;")
class Queryable: class Queryable:
@ -464,7 +489,7 @@ class Queryable:
@property @property
def count(self): def count(self):
return self.db.conn.execute( return self.db.execute(
"select count(*) from [{}]".format(self.name) "select count(*) from [{}]".format(self.name)
).fetchone()[0] ).fetchone()[0]
@ -480,7 +505,7 @@ class Queryable:
sql += " where " + where sql += " where " + where
if order_by is not None: if order_by is not None:
sql += " order by " + order_by 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] columns = [c[0] for c in cursor.description]
for row in cursor: for row in cursor:
yield dict(zip(columns, row)) yield dict(zip(columns, row))
@ -489,9 +514,7 @@ class Queryable:
def columns(self): def columns(self):
if not self.exists(): if not self.exists():
return [] return []
rows = self.db.conn.execute( rows = self.db.execute("PRAGMA table_info([{}])".format(self.name)).fetchall()
"PRAGMA table_info([{}])".format(self.name)
).fetchall()
return [Column(*row) for row in rows] return [Column(*row) for row in rows]
@property @property
@ -501,7 +524,7 @@ class Queryable:
@property @property
def schema(self): def schema(self):
return self.db.conn.execute( return self.db.execute(
"select sql from sqlite_master where name = ?", (self.name,) "select sql from sqlite_master where name = ?", (self.name,)
).fetchone()[0] ).fetchone()[0]
@ -587,7 +610,7 @@ class Table(Queryable):
@property @property
def foreign_keys(self): def foreign_keys(self):
fks = [] fks = []
for row in self.db.conn.execute( for row in self.db.execute(
"PRAGMA foreign_key_list([{}])".format(self.name) "PRAGMA foreign_key_list([{}])".format(self.name)
).fetchall(): ).fetchall():
if row is not None: if row is not None:
@ -615,7 +638,7 @@ class Table(Queryable):
) )
column_sql = "PRAGMA index_info({})".format(index_name_quoted) column_sql = "PRAGMA index_info({})".format(index_name_quoted)
columns = [] 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) columns.append(name)
row["columns"] = columns row["columns"] = columns
# These columns may be missing on older SQLite versions: # These columns may be missing on older SQLite versions:
@ -629,7 +652,7 @@ class Table(Queryable):
def triggers(self): def triggers(self):
return [ return [
Trigger(*r) 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'" "select name, tbl_name, sql from sqlite_master where type = 'trigger'"
" and tbl_name = ?", " and tbl_name = ?",
(self.name,), (self.name,),
@ -683,7 +706,7 @@ class Table(Queryable):
if_not_exists="IF NOT EXISTS " if if_not_exists else "", if_not_exists="IF NOT EXISTS " if if_not_exists else "",
) )
) )
self.db.conn.execute(sql) self.db.execute(sql)
return self return self
def add_column( def add_column(
@ -720,13 +743,13 @@ class Table(Queryable):
col_type=fk_col_type or COLUMN_TYPE_MAPPING[col_type], col_type=fk_col_type or COLUMN_TYPE_MAPPING[col_type],
not_null_default=(" " + not_null_sql) if not_null_sql else "", 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: if fk is not None:
self.add_foreign_key(col_name, fk, fk_col) self.add_foreign_key(col_name, fk, fk_col)
return self return self
def drop(self): 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): def guess_foreign_table(self, column):
column = column.lower() column = column.lower()
@ -811,7 +834,7 @@ class Table(Queryable):
tokenize="\n tokenize='{}',".format(tokenize) if tokenize else "", tokenize="\n tokenize='{}',".format(tokenize) if tokenize else "",
) )
) )
self.db.conn.executescript(sql) self.db.executescript(sql)
self.populate_fts(columns) self.populate_fts(columns)
if create_triggers: if create_triggers:
@ -840,17 +863,23 @@ class Table(Queryable):
new_cols=new_cols, new_cols=new_cols,
) )
) )
self.db.conn.executescript(triggers) self.db.executescript(triggers)
return self return self
def populate_fts(self, columns): def populate_fts(self, columns):
sql = """ sql = (
textwrap.dedent(
"""
INSERT INTO [{table}_fts] (rowid, {columns}) INSERT INTO [{table}_fts] (rowid, {columns})
SELECT rowid, {columns} FROM [{table}]; 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 return self
def disable_fts(self): def disable_fts(self):
@ -858,23 +887,29 @@ class Table(Queryable):
if fts_table: if fts_table:
self.db[fts_table].drop() self.db[fts_table].drop()
# Now delete the triggers that related to that table # Now delete the triggers that related to that table
sql = """ sql = (
textwrap.dedent(
"""
SELECT name FROM sqlite_master SELECT name FROM sqlite_master
WHERE type = 'trigger' WHERE type = 'trigger'
AND sql LIKE '% INSERT INTO [{}]%' AND sql LIKE '% INSERT INTO [{}]%'
""".format( """
fts_table )
.strip()
.format(fts_table)
) )
trigger_names = [] trigger_names = []
for row in self.db.conn.execute(sql).fetchall(): for row in self.db.execute(sql).fetchall():
trigger_names.append(row[0]) trigger_names.append(row[0])
with self.db.conn: with self.db.conn:
for trigger_name in trigger_names: 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): def detect_fts(self):
"Detect if table has a corresponding FTS virtual table and return it" "Detect if table has a corresponding FTS virtual table and return it"
sql = """ sql = (
textwrap.dedent(
"""
SELECT name FROM sqlite_master SELECT name FROM sqlite_master
WHERE rootpage = 0 WHERE rootpage = 0
AND ( AND (
@ -884,10 +919,12 @@ class Table(Queryable):
AND sql LIKE '%VIRTUAL TABLE%USING FTS%' 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: if len(rows) == 0:
return None return None
else: else:
@ -896,18 +933,22 @@ class Table(Queryable):
def optimize(self): def optimize(self):
fts_table = self.detect_fts() fts_table = self.detect_fts()
if fts_table is not None: if fts_table is not None:
self.db.conn.execute( self.db.execute(
""" """
INSERT INTO [{table}] ([{table}]) VALUES ("optimize"); INSERT INTO [{table}] ([{table}]) VALUES ("optimize");
""".format( """.strip().format(
table=fts_table table=fts_table
) )
) )
self.db.conn.execute( self.db.execute(
""" textwrap.dedent(
"""
DELETE FROM [{table}_docsize] WHERE {column} NOT IN ( DELETE FROM [{table}_docsize] WHERE {column} NOT IN (
SELECT rowid FROM [{table}]); SELECT rowid FROM [{table}]);
""".format( """
)
.strip()
.format(
# FTS5 uses 'id' but FTS4 uses 'docid' # FTS5 uses 'id' but FTS4 uses 'docid'
column=self.db["{}_docsize".format(fts_table)].columns[0].name, column=self.db["{}_docsize".format(fts_table)].columns[0].name,
table=fts_table, table=fts_table,
@ -916,16 +957,20 @@ class Table(Queryable):
return self return self
def search(self, q): def search(self, q):
sql = """ sql = (
textwrap.dedent(
"""
select * from "{table}" where rowid in ( select * from "{table}" where rowid in (
select rowid from [{table}_fts] select rowid from [{table}_fts]
where [{table}_fts] match :search where [{table}_fts] match :search
) )
order by rowid 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): def value_or_default(self, key, value):
return self._defaults[key] if value is DEFAULT else 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) table=self.name, wheres=" and ".join(wheres)
) )
with self.db.conn: 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): def delete_where(self, where=None, where_args=None):
if not self.exists(): if not self.exists():
@ -947,7 +992,7 @@ class Table(Queryable):
sql = "delete from [{}]".format(self.name) sql = "delete from [{}]".format(self.name)
if where is not None: if where is not None:
sql += " where " + where 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): def update(self, pk_values, updates=None, alter=False, conversions=None):
updates = updates or {} updates = updates or {}
@ -972,12 +1017,12 @@ class Table(Queryable):
) )
with self.db.conn: with self.db.conn:
try: try:
rowcount = self.db.conn.execute(sql, args).rowcount rowcount = self.db.execute(sql, args).rowcount
except OperationalError as e: except OperationalError as e:
if alter and (" column" in e.args[0]): if alter and (" column" in e.args[0]):
# Attempt to add any missing columns, then try again # Attempt to add any missing columns, then try again
self.add_missing_columns([updates]) self.add_missing_columns([updates])
rowcount = self.db.conn.execute(sql, args).rowcount rowcount = self.db.execute(sql, args).rowcount
else: else:
raise raise
@ -1084,7 +1129,7 @@ class Table(Queryable):
self.last_rowid = None self.last_rowid = None
self.last_pk = None self.last_pk = None
if truncate and self.exists(): 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): for chunk in chunks(itertools.chain([first_record], records), batch_size):
chunk = list(chunk) chunk = list(chunk)
num_records_processed += len(chunk) num_records_processed += len(chunk)
@ -1184,14 +1229,12 @@ class Table(Queryable):
or_what = "OR IGNORE " or_what = "OR IGNORE "
sql = """ sql = """
INSERT {or_what}INTO [{table}] ({columns}) VALUES {rows}; INSERT {or_what}INTO [{table}] ({columns}) VALUES {rows};
""".format( """.strip().format(
or_what=or_what, or_what=or_what,
table=self.name, table=self.name,
columns=", ".join("[{}]".format(c) for c in all_columns), columns=", ".join("[{}]".format(c) for c in all_columns),
rows=", ".join( rows=", ".join(
""" "({placeholders})".format(
({placeholders})
""".format(
placeholders=", ".join( placeholders=", ".join(
[conversions.get(col, "?") for col in all_columns] [conversions.get(col, "?") for col in all_columns]
) )
@ -1205,12 +1248,12 @@ class Table(Queryable):
with self.db.conn: with self.db.conn:
for query, params in queries_and_params: for query, params in queries_and_params:
try: try:
result = self.db.conn.execute(query, params) result = self.db.execute(query, params)
except OperationalError as e: except OperationalError as e:
if alter and (" column" in e.args[0]): if alter and (" column" in e.args[0]):
# Attempt to add any missing columns, then try again # Attempt to add any missing columns, then try again
self.add_missing_columns(chunk) self.add_missing_columns(chunk)
result = self.db.conn.execute(query, params) result = self.db.execute(query, params)
else: else:
raise raise
if num_records_processed == 1 and not upsert: if num_records_processed == 1 and not upsert:
@ -1383,7 +1426,7 @@ class View(Queryable):
) )
def drop(self): def drop(self):
self.db.conn.execute("DROP VIEW [{}]".format(self.name)) self.db.execute("DROP VIEW [{}]".format(self.name))
def chunks(sequence, size): def chunks(sequence, size):

View file

@ -10,7 +10,7 @@ def fresh_db():
@pytest.fixture @pytest.fixture
def existing_db(): def existing_db():
database = Database(memory=True) database = Database(memory=True)
database.conn.executescript( database.executescript(
""" """
CREATE TABLE foo (text TEXT); CREATE TABLE foo (text TEXT);
INSERT INTO foo (text) values ("one"); INSERT INTO foo (text) values ("one");

View file

@ -396,7 +396,7 @@ def test_enable_fts_with_triggers(db_path):
def search(q): def search(q):
return ( return (
Database(db_path) 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() .fetchall()
) )
@ -417,7 +417,7 @@ def test_populate_fts(db_path):
def search(q): def search(q):
return ( return (
Database(db_path) 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() .fetchall()
) )

View file

@ -41,5 +41,5 @@ def test_column_affinity(column_def, expected_type):
@pytest.mark.parametrize("column_def,expected_type", EXAMPLES) @pytest.mark.parametrize("column_def,expected_type", EXAMPLES)
def test_columns_dict(fresh_db, column_def, expected_type): 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 assert {"col": expected_type} == fresh_db["foo"].columns_dict

View file

@ -4,9 +4,9 @@ import pytest
def test_recursive_triggers(): def test_recursive_triggers():
db = Database(memory=True) 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(): def test_recursive_triggers_off():
db = Database(memory=True, recursive_triggers=False) 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]

View file

@ -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): def test_insert_dictionaries_and_lists_as_json(fresh_db, data_structure):
fresh_db["test"].insert({"id": 1, "data": data_structure}, pk="id") 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 row[0] == 1
assert data_structure == json.loads(row[1]) assert data_structure == json.loads(row[1])

View file

@ -4,7 +4,7 @@ from sqlite_utils.utils import OperationalError
def test_create_view(fresh_db): def test_create_view(fresh_db):
fresh_db.create_view("bar", "select 1 + 1") 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 assert [(2,)] == rows
@ -23,7 +23,7 @@ def test_create_view_ignore(fresh_db):
fresh_db.create_view("bar", "select 1 + 1").create_view( fresh_db.create_view("bar", "select 1 + 1").create_view(
"bar", "select 1 + 2", ignore=True "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 assert [(2,)] == rows
@ -31,13 +31,13 @@ def test_create_view_replace(fresh_db):
fresh_db.create_view("bar", "select 1 + 1").create_view( fresh_db.create_view("bar", "select 1 + 1").create_view(
"bar", "select 1 + 2", replace=True "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 assert [(3,)] == rows
def test_create_view_replace_with_same_does_nothing(fresh_db): def test_create_view_replace_with_same_does_nothing(fresh_db):
fresh_db.create_view("bar", "select 1 + 1") 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) 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 assert after_version == initial_version

View file

@ -160,7 +160,7 @@ def test_disable_fts(fresh_db, create_triggers):
expected_triggers = set() expected_triggers = set()
assert expected_triggers == set( assert expected_triggers == set(
r[0] r[0]
for r in fresh_db.conn.execute( for r in fresh_db.execute(
"select name from sqlite_master where type = 'trigger'" "select name from sqlite_master where type = 'trigger'"
).fetchall() ).fetchall()
) )
@ -168,7 +168,7 @@ def test_disable_fts(fresh_db, create_triggers):
table.disable_fts() table.disable_fts()
assert ( assert (
0 0
== fresh_db.conn.execute( == fresh_db.execute(
"select count(*) from sqlite_master where type = 'trigger'" "select count(*) from sqlite_master where type = 'trigger'"
).fetchone()[0] ).fetchone()[0]
) )

View file

@ -73,7 +73,7 @@ def test_table_repr(fresh_db):
def test_indexes(fresh_db): def test_indexes(fresh_db):
fresh_db.conn.executescript( fresh_db.executescript(
""" """
create table Gosh (c1 text, c2 text, c3 text); create table Gosh (c1 text, c2 text, c3 text);
create index Gosh_c1 on Gosh(c1); create index Gosh_c1 on Gosh(c1);

62
tests/test_tracer.py Normal file
View file

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