Sketched out a trace mechanism, refs #150

This commit is contained in:
Simon Willison 2020-09-07 12:44:37 -07:00
commit 6104cbd91a

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
@ -124,6 +126,16 @@ class Database:
self.conn = filename_or_conn self.conn = filename_or_conn
if recursive_triggers: if recursive_triggers:
self.conn.execute("PRAGMA recursive_triggers=on;") self.conn.execute("PRAGMA recursive_triggers=on;")
self._tracer = tracer
@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)
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,7 +863,7 @@ 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):
@ -850,7 +873,7 @@ class Table(Queryable):
""".format( """.format(
table=self.name, columns=", ".join("[{}]".format(c) for c in columns) 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):
@ -866,11 +889,11 @@ class Table(Queryable):
fts_table 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"
@ -887,7 +910,7 @@ class Table(Queryable):
""".format( """.format(
table=self.name 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,7 +919,7 @@ 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( """.format(
@ -925,7 +948,7 @@ class Table(Queryable):
""".format( """.format(
table=self.name 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 +962,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 +970,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 +995,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 +1107,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)
@ -1205,12 +1228,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 +1406,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):