mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-09-28 04:44:26 +02:00
Fixed PRAGMA foreign_keys handling for .transform, closes #167
This commit is contained in:
parent
5c4d58d152
commit
b8e0048485
2 changed files with 22 additions and 26 deletions
|
|
@ -741,20 +741,21 @@ class Table(Queryable):
|
||||||
defaults=defaults,
|
defaults=defaults,
|
||||||
drop_foreign_keys=drop_foreign_keys,
|
drop_foreign_keys=drop_foreign_keys,
|
||||||
)
|
)
|
||||||
initial_pragma_foreign_keys = self.db.execute("PRAGMA foreign_keys").fetchone()[
|
pragma_foreign_keys_was_on = self.db.execute(
|
||||||
0
|
"PRAGMA foreign_keys"
|
||||||
]
|
).fetchone()[0]
|
||||||
try:
|
try:
|
||||||
|
if pragma_foreign_keys_was_on:
|
||||||
|
self.db.execute("PRAGMA foreign_keys=0;")
|
||||||
with self.db.conn:
|
with self.db.conn:
|
||||||
for sql in sqls:
|
for sql in sqls:
|
||||||
self.db.execute(sql)
|
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:
|
finally:
|
||||||
# Make sure we reset PRAGMA foreign_keys correctly
|
if pragma_foreign_keys_was_on:
|
||||||
if (
|
self.db.execute("PRAGMA foreign_keys=1;")
|
||||||
initial_pragma_foreign_keys
|
|
||||||
and not self.db.execute("PRAGMA foreign_keys").fetchone()[0]
|
|
||||||
):
|
|
||||||
self.db.execute("PRAGMA foreign_keys=1")
|
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def transform_sql(
|
def transform_sql(
|
||||||
|
|
@ -787,15 +788,7 @@ class Table(Queryable):
|
||||||
new_column_pairs.append((new_name, type_))
|
new_column_pairs.append((new_name, type_))
|
||||||
copy_from_to[name] = new_name
|
copy_from_to[name] = new_name
|
||||||
|
|
||||||
should_flip_foreign_keys_pragma = self.db.execute(
|
|
||||||
"PRAGMA foreign_keys"
|
|
||||||
).fetchone()[0]
|
|
||||||
|
|
||||||
sqls = []
|
sqls = []
|
||||||
|
|
||||||
if should_flip_foreign_keys_pragma:
|
|
||||||
sqls.append("PRAGMA foreign_keys=OFF;")
|
|
||||||
|
|
||||||
if pk is DEFAULT:
|
if pk is DEFAULT:
|
||||||
pks_renamed = tuple(rename.get(p) or p for p in self.pks)
|
pks_renamed = tuple(rename.get(p) or p for p in self.pks)
|
||||||
if len(pks_renamed) == 1:
|
if len(pks_renamed) == 1:
|
||||||
|
|
@ -876,11 +869,6 @@ class Table(Queryable):
|
||||||
sqls.append(
|
sqls.append(
|
||||||
"ALTER TABLE [{}] RENAME TO [{}];".format(new_table_name, self.name)
|
"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
|
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, progress=None):
|
||||||
|
|
|
||||||
|
|
@ -90,17 +90,25 @@ import pytest
|
||||||
)
|
)
|
||||||
@pytest.mark.parametrize("use_pragma_foreign_keys", [False, True])
|
@pytest.mark.parametrize("use_pragma_foreign_keys", [False, True])
|
||||||
def test_transform_sql(fresh_db, params, expected_sql, use_pragma_foreign_keys):
|
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"]
|
dogs = fresh_db["dogs"]
|
||||||
if use_pragma_foreign_keys:
|
if use_pragma_foreign_keys:
|
||||||
fresh_db.conn.execute("PRAGMA foreign_keys=ON")
|
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")
|
dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id")
|
||||||
sql = dogs.transform_sql(**{**params, **{"tmp_suffix": "suffix"}})
|
sql = dogs.transform_sql(**{**params, **{"tmp_suffix": "suffix"}})
|
||||||
assert sql == expected_sql
|
assert sql == expected_sql
|
||||||
# Check that .transform() runs without exceptions:
|
# 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):
|
def test_transform_sql_rowid_to_id(fresh_db):
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue