diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index e9d5298..cbd657d 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -741,20 +741,21 @@ class Table(Queryable): defaults=defaults, drop_foreign_keys=drop_foreign_keys, ) - initial_pragma_foreign_keys = self.db.execute("PRAGMA foreign_keys").fetchone()[ - 0 - ] + pragma_foreign_keys_was_on = self.db.execute( + "PRAGMA foreign_keys" + ).fetchone()[0] try: + if pragma_foreign_keys_was_on: + self.db.execute("PRAGMA foreign_keys=0;") with self.db.conn: for sql in sqls: self.db.execute(sql) + # Run the foreign_key_check before we commit + if pragma_foreign_keys_was_on: + self.db.execute("PRAGMA foreign_key_check;") 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") + if pragma_foreign_keys_was_on: + self.db.execute("PRAGMA foreign_keys=1;") return self def transform_sql( @@ -787,15 +788,7 @@ 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: pks_renamed = tuple(rename.get(p) or p for p in self.pks) if len(pks_renamed) == 1: @@ -876,11 +869,6 @@ class Table(Queryable): 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 extract(self, columns, table=None, fk_column=None, rename=None, progress=None): diff --git a/tests/test_transform.py b/tests/test_transform.py index 352efe4..7075fcf 100644 --- a/tests/test_transform.py +++ b/tests/test_transform.py @@ -90,17 +90,25 @@ import pytest ) @pytest.mark.parametrize("use_pragma_foreign_keys", [False, True]) def test_transform_sql(fresh_db, params, expected_sql, use_pragma_foreign_keys): + captured = [] + tracer = lambda sql, params: captured.append((sql, params)) 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 # Check that .transform() runs without exceptions: - dogs.transform(**params) + with fresh_db.tracer(tracer): + dogs.transform(**params) + # If use_pragma_foreign_keys, check that we did the right thing + if use_pragma_foreign_keys: + assert ('PRAGMA foreign_keys=0;', None) in captured + assert captured[-2] == ('PRAGMA foreign_key_check;', None) + assert captured[-1] == ('PRAGMA foreign_keys=1;', None) + else: + assert ('PRAGMA foreign_keys=0;', None) not in captured + assert ('PRAGMA foreign_keys=1;', None) not in captured def test_transform_sql_rowid_to_id(fresh_db):