From dc3eb9c313502832c07f4de889aa60c373f45b52 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Mon, 21 Sep 2020 21:18:24 -0700 Subject: [PATCH] Correct handling of pragma foreign_keys for .transform() --- sqlite_utils/db.py | 37 +++++++++++++++++++++++++----- tests/test_transform.py | 51 ++++++++++++++++++++++++++++++++++++++--- 2 files changed, 79 insertions(+), 9 deletions(-) diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 02b55ec..d8ad05b 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -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): diff --git a/tests/test_transform.py b/tests/test_transform.py index 2fbaa6b..c49efe1 100644 --- a/tests/test_transform.py +++ b/tests/test_transform.py @@ -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]