Correct handling of pragma foreign_keys for .transform()

This commit is contained in:
Simon Willison 2020-09-21 21:18:24 -07:00
commit dc3eb9c313
2 changed files with 79 additions and 9 deletions

View file

@ -736,9 +736,20 @@ class Table(Queryable):
defaults=defaults,
drop_foreign_keys=drop_foreign_keys,
)
with self.db.conn:
for sql in sqls:
self.db.conn.execute(sql)
initial_pragma_foreign_keys = self.db.execute("PRAGMA foreign_keys").fetchone()[
0
]
try:
with self.db.conn:
for sql in sqls:
self.db.execute(sql)
finally:
# Make sure we reset PRAGMA foreign_keys correctly
if (
initial_pragma_foreign_keys
and not self.db.execute("PRAGMA foreign_keys").fetchone()[0]
):
self.db.execute("PRAGMA foreign_keys=1")
return self
def transform_sql(
@ -770,12 +781,21 @@ class Table(Queryable):
new_column_pairs.append((new_name, type_))
copy_from_to[name] = new_name
should_flip_foreign_keys_pragma = self.db.execute(
"PRAGMA foreign_keys"
).fetchone()[0]
sqls = []
if should_flip_foreign_keys_pragma:
sqls.append("PRAGMA foreign_keys=OFF")
if pk is DEFAULT:
if len(self.pks) == 1:
pk = self.pks[0]
pks_renamed = tuple(rename.get(p) or p for p in self.pks)
if len(pks_renamed) == 1:
pk = pks_renamed[0]
else:
pk = self.pks
pk = pks_renamed
# not_null may be a set or dict, need to convert to a set
create_table_not_null = {c.name for c in self.columns if c.notnull}
@ -839,6 +859,11 @@ class Table(Queryable):
sqls.append("DROP TABLE [{}]".format(self.name))
# Rename the new one
sqls.append("ALTER TABLE [{}] RENAME TO [{}]".format(new_table_name, self.name))
if should_flip_foreign_keys_pragma:
sqls.append("PRAGMA foreign_key_check")
sqls.append("PRAGMA foreign_keys=ON")
return sqls
def create_index(self, columns, index_name=None, unique=False, if_not_exists=False):

View file

@ -1,4 +1,5 @@
from sqlite_utils.db import ForeignKey
from sqlite_utils.utils import OperationalError
import pytest
@ -87,8 +88,14 @@ import pytest
),
],
)
def test_transform_sql(fresh_db, params, expected_sql):
@pytest.mark.parametrize("use_pragma_foreign_keys", [False, True])
def test_transform_sql(fresh_db, params, expected_sql, use_pragma_foreign_keys):
dogs = fresh_db["dogs"]
if use_pragma_foreign_keys:
fresh_db.conn.execute("PRAGMA foreign_keys=ON")
expected_sql.insert(0, "PRAGMA foreign_keys=OFF")
expected_sql.append("PRAGMA foreign_key_check")
expected_sql.append("PRAGMA foreign_keys=ON")
dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id")
sql = dogs.transform_sql(**{**params, **{"tmp_suffix": "suffix"}})
assert sql == expected_sql
@ -111,6 +118,16 @@ def test_transform_sql_rowid_to_id(fresh_db):
)
def test_transform_rename_pk(fresh_db):
dogs = fresh_db["dogs"]
dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id")
dogs.transform(rename={"id": "pk"})
assert (
dogs.schema
== 'CREATE TABLE "dogs" (\n [pk] INTEGER PRIMARY KEY,\n [name] TEXT,\n [age] TEXT\n)'
)
def test_transform_not_null(fresh_db):
dogs = fresh_db["dogs"]
dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id")
@ -199,7 +216,12 @@ def test_transform_foreign_keys_persist(authors_db):
]
def test_transform_foreign_keys_survive_renamed_column(authors_db):
@pytest.mark.parametrize("use_pragma_foreign_keys", [False, True])
def test_transform_foreign_keys_survive_renamed_column(
authors_db, use_pragma_foreign_keys
):
if use_pragma_foreign_keys:
authors_db.conn.execute("PRAGMA foreign_keys=ON")
authors_db["books"].transform(rename={"author_id": "author_id_2"})
assert authors_db["books"].foreign_keys == [
ForeignKey(
@ -211,7 +233,10 @@ def test_transform_foreign_keys_survive_renamed_column(authors_db):
]
def test_transform_drop_foreign_keys(fresh_db):
@pytest.mark.parametrize("use_pragma_foreign_keys", [False, True])
def test_transform_drop_foreign_keys(fresh_db, use_pragma_foreign_keys):
if use_pragma_foreign_keys:
fresh_db.conn.execute("PRAGMA foreign_keys=ON")
# Create table with three foreign keys so we can drop two of them
fresh_db["country"].insert({"id": 1, "name": "France"}, pk="id")
fresh_db["continent"].insert({"id": 2, "name": "Europe"}, pk="id")
@ -251,3 +276,23 @@ def test_transform_drop_foreign_keys(fresh_db):
assert fresh_db["places"].foreign_keys == [
ForeignKey(table="places", column="city", other_table="city", other_column="id")
]
if use_pragma_foreign_keys:
assert fresh_db.conn.execute("PRAGMA foreign_keys").fetchone()[0]
def test_transform_verify_foreign_keys(fresh_db):
fresh_db.conn.execute("PRAGMA foreign_keys=ON")
fresh_db["authors"].insert({"id": 3, "name": "Tina"}, pk="id")
fresh_db["books"].insert(
{"id": 1, "title": "Book", "author_id": 3}, pk="id", foreign_keys={"author_id"}
)
# Renaming the id column on authors should break everything
with pytest.raises(OperationalError) as e:
fresh_db["authors"].transform(rename={"id": "id2"})
assert e.value.args[0] == 'foreign key mismatch - "books" referencing "authors"'
# This should have rolled us back
assert (
fresh_db["authors"].schema
== "CREATE TABLE [authors] (\n [id] INTEGER PRIMARY KEY,\n [name] TEXT\n)"
)
assert fresh_db.conn.execute("PRAGMA foreign_keys").fetchone()[0]