Much, much faster extract() implementation

Takes my test down from ten minutes to four seconds!

* Removed unnecessary update() optimization
* Added column_order= to .transform() and .transform_sql()
* Tests for reusing lookup table in extract()

Closes #172
This commit is contained in:
Simon Willison 2020-09-24 08:43:55 -07:00 • committed by GitHub
commit 022cdd97a9
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 135 additions and 58 deletions

View file

@ -1006,12 +1006,6 @@ def transform(
multiple=True, multiple=True,
help="Rename this column in extracted table", help="Rename this column in extracted table",
) )
@click.option(
"-s",
"--silent",
is_flag=True,
help="Don't show progress bar",
)
def extract( def extract(
path, path,
table, table,
@ -1019,7 +1013,6 @@ def extract(
other_table, other_table,
fk_column, fk_column,
rename, rename,
silent,
): ):
"Extract one or more columns into a separate table" "Extract one or more columns into a separate table"
db = sqlite_utils.Database(path) db = sqlite_utils.Database(path)
@ -1029,14 +1022,7 @@ def extract(
fk_column=fk_column, fk_column=fk_column,
rename=dict(rename), rename=dict(rename),
) )
if silent: db[table].extract(**kwargs)
db[table].extract(**kwargs)
else:
with click.progressbar(
length=db[table].count, label="Extracting columns"
) as bar:
kwargs["progress"] = bar.update
db[table].extract(**kwargs)
@cli.command(name="insert-files") @cli.command(name="insert-files")

View file

@ -3,7 +3,6 @@ from collections import namedtuple, OrderedDict
import contextlib import contextlib
import datetime import datetime
import decimal import decimal
import functools
import hashlib import hashlib
import inspect import inspect
import itertools import itertools
@ -731,6 +730,7 @@ class Table(Queryable):
not_null=None, not_null=None,
defaults=None, defaults=None,
drop_foreign_keys=None, drop_foreign_keys=None,
column_order=None,
): ):
assert self.exists(), "Cannot transform a table that doesn't exist yet" assert self.exists(), "Cannot transform a table that doesn't exist yet"
sqls = self.transform_sql( sqls = self.transform_sql(
@ -741,6 +741,7 @@ class Table(Queryable):
not_null=not_null, not_null=not_null,
defaults=defaults, defaults=defaults,
drop_foreign_keys=drop_foreign_keys, drop_foreign_keys=drop_foreign_keys,
column_order=column_order,
) )
pragma_foreign_keys_was_on = self.db.execute("PRAGMA foreign_keys").fetchone()[ pragma_foreign_keys_was_on = self.db.execute("PRAGMA foreign_keys").fetchone()[
0 0
@ -769,6 +770,7 @@ class Table(Queryable):
not_null=None, not_null=None,
defaults=None, defaults=None,
drop_foreign_keys=None, drop_foreign_keys=None,
column_order=None,
tmp_suffix=None, tmp_suffix=None,
): ):
types = types or {} types = types or {}
@ -840,6 +842,9 @@ class Table(Queryable):
(rename.get(column) or column, other_table, other_column) (rename.get(column) or column, other_table, other_column)
) )
if column_order is not None:
column_order = [rename.get(col) or col for col in column_order]
sqls.append( sqls.append(
self.db.create_table_sql( self.db.create_table_sql(
new_table_name, new_table_name,
@ -848,6 +853,7 @@ class Table(Queryable):
not_null=create_table_not_null, not_null=create_table_not_null,
defaults=create_table_defaults, defaults=create_table_defaults,
foreign_keys=create_table_foreign_keys, foreign_keys=create_table_foreign_keys,
column_order=column_order,
).strip() ).strip()
) )
@ -872,7 +878,7 @@ class Table(Queryable):
) )
return sqls return sqls
def extract(self, columns, table=None, fk_column=None, rename=None, progress=None): def extract(self, columns, table=None, fk_column=None, rename=None):
rename = rename or {} rename = rename or {}
if isinstance(columns, str): if isinstance(columns, str):
columns = [columns] columns = [columns]
@ -886,38 +892,85 @@ class Table(Queryable):
first_column = columns[0] first_column = columns[0]
pks = self.pks pks = self.pks
lookup_table = self.db[table] lookup_table = self.db[table]
@functools.lru_cache(maxsize=128)
def cached_lookup(lookups):
return lookup_table.lookup(dict(lookups))
if pks == ["rowid"]:
rows_iter = self.rows_where(select="rowid, *")
else:
rows_iter = self.rows
for row in rows_iter:
row_pks = tuple(row[pk] for pk in pks)
lookups = tuple(
(rename.get(column) or column, row[column]) for column in columns
)
self.update(
row_pks,
{first_column: cached_lookup(lookups)},
assume_exists=True,
pks=self.pks,
)
if progress:
progress(1)
fk_column = fk_column or "{}_id".format(table) fk_column = fk_column or "{}_id".format(table)
# Now rename first_column and change its type to integer, and drop magic_lookup_column = "{}_{}".format(fk_column, os.urandom(6).hex())
# any other extracted columns:
self.transform( # Populate the lookup table with all of the extracted unique values
types={first_column: int}, lookup_columns_definition = {
drop=set(columns[1:]), (rename.get(col) or col): typ
rename={first_column: fk_column}, for col, typ in self.columns_dict.items()
if col in columns
}
if lookup_table.exists():
if not set(lookup_columns_definition.items()).issubset(
lookup_table.columns_dict.items()
):
raise InvalidColumns(
"Lookup table {} already exists but does not have columns {}".format(
table, lookup_columns_definition
)
)
else:
lookup_table.create(
{
**{
"id": int,
},
**lookup_columns_definition,
},
pk="id",
)
lookup_columns = [(rename.get(col) or col) for col in columns]
lookup_table.create_index(lookup_columns, unique=True, if_not_exists=True)
self.db.execute(
"INSERT OR IGNORE INTO [{lookup_table}] ({lookup_columns}) SELECT DISTINCT {table_cols} FROM [{table}]".format(
lookup_table=table,
lookup_columns=", ".join("[{}]".format(c) for c in lookup_columns),
table_cols=", ".join("[{}]".format(c) for c in columns),
table=self.name,
)
) )
# Now add the new fk_column
self.add_column(magic_lookup_column, int)
# And populate it
self.db.execute(
"UPDATE [{table}] SET [{magic_lookup_column}] = (SELECT id FROM [{lookup_table}] WHERE {where})".format(
table=self.name,
magic_lookup_column=magic_lookup_column,
lookup_table=table,
where=" AND ".join(
"[{table}].[{column}] = [{lookup_table}].[{lookup_column}]".format(
table=self.name,
lookup_table=table,
column=column,
lookup_column=rename.get(column) or column,
)
for column in columns
),
)
)
# Figure out the right column order
column_order = []
for c in self.columns:
if c.name in columns and magic_lookup_column not in column_order:
column_order.append(magic_lookup_column)
elif c.name == magic_lookup_column:
continue
else:
column_order.append(c.name)
# Drop the unnecessary columns and rename lookup column
self.transform(
drop=set(columns),
rename={magic_lookup_column: fk_column},
column_order=column_order,
)
# And add the foreign key constraint # And add the foreign key constraint
self.add_foreign_key(fk_column, table, "id") self.add_foreign_key(fk_column, table, "id")
return self
def create_index(self, columns, index_name=None, unique=False, if_not_exists=False): def create_index(self, columns, index_name=None, unique=False, if_not_exists=False):
if index_name is None: if index_name is None:
@ -1258,28 +1311,19 @@ class Table(Queryable):
self.db.execute(sql, where_args or []) self.db.execute(sql, where_args or [])
return self return self
def update( def update(self, pk_values, updates=None, alter=False, conversions=None):
self,
pk_values,
updates=None,
alter=False,
conversions=None,
assume_exists=False,
pks=None,
):
updates = updates or {} updates = updates or {}
pks = pks or self.pks
conversions = conversions or {} conversions = conversions or {}
if not isinstance(pk_values, (list, tuple)): if not isinstance(pk_values, (list, tuple)):
pk_values = [pk_values] pk_values = [pk_values]
# Soundness check that the record exists (raises error if not): # Soundness check that the record exists (raises error if not):
if not assume_exists: self.get(pk_values)
self.get(pk_values)
if not updates: if not updates:
return self return self
args = [] args = []
sets = [] sets = []
wheres = [] wheres = []
pks = self.pks
validate_column_names(updates.keys()) validate_column_names(updates.keys())
for key, value in updates.items(): for key, value in updates.items():
sets.append("[{}] = {}".format(key, conversions.get(key, "?"))) sets.append("[{}] = {}".format(key, conversions.get(key, "?")))

View file

@ -11,7 +11,12 @@ def test_extract_single_column(fresh_db, table, fk_column):
iter_species = itertools.cycle(["Palm", "Spruce", "Mangrove", "Oak"]) iter_species = itertools.cycle(["Palm", "Spruce", "Mangrove", "Oak"])
fresh_db["tree"].insert_all( fresh_db["tree"].insert_all(
( (
{"id": i, "name": "Tree {}".format(i), "species": next(iter_species)} {
"id": i,
"name": "Tree {}".format(i),
"species": next(iter_species),
"end": 1,
}
for i in range(1, 1001) for i in range(1, 1001)
), ),
pk="id", pk="id",
@ -22,6 +27,7 @@ def test_extract_single_column(fresh_db, table, fk_column):
" [id] INTEGER PRIMARY KEY,\n" " [id] INTEGER PRIMARY KEY,\n"
" [name] TEXT,\n" " [name] TEXT,\n"
" [{}] INTEGER,\n".format(expected_fk) " [{}] INTEGER,\n".format(expected_fk)
+ " [end] INTEGER,\n"
+ " FOREIGN KEY({}) REFERENCES {}(id)\n".format(expected_fk, expected_table) + " FOREIGN KEY({}) REFERENCES {}(id)\n".format(expected_fk, expected_table)
+ ")" + ")"
) )
@ -38,10 +44,10 @@ def test_extract_single_column(fresh_db, table, fk_column):
{"id": 4, "species": "Oak"}, {"id": 4, "species": "Oak"},
] ]
assert list(itertools.islice(fresh_db["tree"].rows, 0, 4)) == [ assert list(itertools.islice(fresh_db["tree"].rows, 0, 4)) == [
{"id": 1, "name": "Tree 1", expected_fk: 1}, {"id": 1, "name": "Tree 1", expected_fk: 1, "end": 1},
{"id": 2, "name": "Tree 2", expected_fk: 2}, {"id": 2, "name": "Tree 2", expected_fk: 2, "end": 1},
{"id": 3, "name": "Tree 3", expected_fk: 3}, {"id": 3, "name": "Tree 3", expected_fk: 3, "end": 1},
{"id": 4, "name": "Tree 4", expected_fk: 4}, {"id": 4, "name": "Tree 4", expected_fk: 4, "end": 1},
] ]
@ -124,3 +130,44 @@ def test_extract_rowid_table(fresh_db):
" FOREIGN KEY(common_name_latin_name_id) REFERENCES common_name_latin_name(id)\n" " FOREIGN KEY(common_name_latin_name_id) REFERENCES common_name_latin_name(id)\n"
")" ")"
) )
def test_reuse_lookup_table(fresh_db):
fresh_db["species"].insert({"id": 1, "name": "Wolf"}, pk="id")
fresh_db["sightings"].insert({"id": 10, "species": "Wolf"}, pk="id")
fresh_db["individuals"].insert(
{"id": 10, "name": "Terriana", "species": "Fox"}, pk="id"
)
fresh_db["sightings"].extract("species", rename={"species": "name"})
fresh_db["individuals"].extract("species", rename={"species": "name"})
assert fresh_db["sightings"].schema == (
'CREATE TABLE "sightings" (\n'
" [id] INTEGER PRIMARY KEY,\n"
" [species_id] INTEGER,\n"
" FOREIGN KEY(species_id) REFERENCES species(id)\n"
")"
)
assert fresh_db["individuals"].schema == (
'CREATE TABLE "individuals" (\n'
" [id] INTEGER PRIMARY KEY,\n"
" [name] TEXT,\n"
" [species_id] INTEGER,\n"
" FOREIGN KEY(species_id) REFERENCES species(id)\n"
")"
)
assert list(fresh_db["species"].rows) == [
{"id": 1, "name": "Wolf"},
{"id": 2, "name": "Fox"},
]
def test_extract_error_on_incompatible_existing_lookup_table(fresh_db):
fresh_db["species"].insert({"id": 1})
fresh_db["tree"].insert({"name": "Tree 1", "common_name": "Palm"})
with pytest.raises(InvalidColumns):
fresh_db["tree"].extract("common_name", table="species")
# Try again with incompatible existing column type
fresh_db["species2"].insert({"id": 1, "common_name": 3.5})
with pytest.raises(InvalidColumns):
fresh_db["tree"].extract("common_name", table="species2")