.transform() preserves check constraints, refs #762

This commit is contained in:
Simon Willison 2026-08-11 21:59:57 -07:00
commit 2303b80aef
6 changed files with 447 additions and 2 deletions

View file

@ -0,0 +1,154 @@
import pytest
from sqlite_utils import Database
from sqlite_utils.utils import sqlite3
BASELINE_ROWS = [(1, "one"), (2, "two")]
def insert(table):
table.insert({"id": 3, "value": "three"}, pk="id")
def insert_all(table):
table.insert_all(
[
{"id": 3, "value": "three"},
{"id": 4, "value": "four"},
],
pk="id",
batch_size=1,
)
def upsert(table):
table.upsert({"id": 2, "value": "TWO"}, pk="id")
def upsert_all(table):
table.upsert_all(
[
{"id": 2, "value": "TWO"},
{"id": 3, "value": "three"},
],
pk="id",
batch_size=1,
)
def update(table):
table.update(2, {"value": "TWO"})
def delete(table):
table.delete(2)
def delete_where(table):
table.delete_where("id > ?", [1])
MUTATOR_CASES = (
pytest.param(
insert,
[(1, "one"), (2, "two"), (3, "three")],
id="insert",
),
pytest.param(
insert_all,
[(1, "one"), (2, "two"), (3, "three"), (4, "four")],
id="insert_all",
),
pytest.param(
upsert,
[(1, "one"), (2, "TWO")],
id="upsert",
),
pytest.param(
upsert_all,
[(1, "one"), (2, "TWO"), (3, "three")],
id="upsert_all",
),
pytest.param(
update,
[(1, "one"), (2, "TWO")],
id="update",
),
pytest.param(delete, [(1, "one")], id="delete"),
pytest.param(delete_where, [(1, "one")], id="delete_where"),
)
class RollbackTest(Exception):
pass
def seed_database(path):
conn = sqlite3.connect(str(path))
try:
conn.execute("create table items (id integer primary key, value text)")
conn.executemany("insert into items values (?, ?)", BASELINE_ROWS)
conn.commit()
finally:
conn.close()
return Database(path)
def current_rows(db):
return db.conn.execute("select id, value from items order by id").fetchall()
def persisted_rows(path):
conn = sqlite3.connect(str(path))
try:
return conn.execute("select id, value from items order by id").fetchall()
finally:
conn.close()
@pytest.mark.parametrize("mutate,expected_rows", MUTATOR_CASES)
def test_mutator_commits_by_default(tmp_path, mutate, expected_rows):
path = tmp_path / "default.db"
db = seed_database(path)
assert not db.conn.in_transaction
mutate(db["items"])
assert current_rows(db) == expected_rows
assert not db.conn.in_transaction
db.close()
assert persisted_rows(path) == expected_rows
@pytest.mark.parametrize("mutate,expected_rows", MUTATOR_CASES)
def test_mutator_commits_with_outer_atomic(tmp_path, mutate, expected_rows):
path = tmp_path / "atomic.db"
db = seed_database(path)
with db.atomic():
assert db.conn.in_transaction
mutate(db["items"])
assert current_rows(db) == expected_rows
assert db.conn.in_transaction
assert current_rows(db) == expected_rows
assert not db.conn.in_transaction
db.close()
assert persisted_rows(path) == expected_rows
@pytest.mark.parametrize("mutate,expected_rows", MUTATOR_CASES)
def test_mutator_rolls_back_outer_atomic(tmp_path, mutate, expected_rows):
path = tmp_path / "rollback.db"
db = seed_database(path)
with pytest.raises(RollbackTest), db.atomic():
mutate(db["items"])
assert current_rows(db) == expected_rows
assert db.conn.in_transaction
raise RollbackTest
assert current_rows(db) == BASELINE_ROWS
assert not db.conn.in_transaction
db.close()
assert persisted_rows(path) == BASELINE_ROWS

View file

@ -2,7 +2,7 @@ import sqlite3
import pytest
from sqlite_utils.db import ForeignKey, TransactionError, TransformError
from sqlite_utils.db import Check, ForeignKey, TransactionError, TransformError
from sqlite_utils.utils import OperationalError
@ -1065,3 +1065,138 @@ def test_transform_restores_legacy_alter_table_setting(fresh_db):
assert sqls[-1] == "PRAGMA legacy_alter_table=ON;"
dogs.transform(types={"name": str})
assert fresh_db.execute("PRAGMA legacy_alter_table").fetchone()[0] == 1
def test_transform_preserves_check_constraints(fresh_db):
fresh_db.execute("""
CREATE TABLE scores (
id INTEGER PRIMARY KEY,
score INTEGER CONSTRAINT valid_score CHECK(score BETWEEN 0 AND 100),
CONSTRAINT nonzero_id CHECK(id != 0)
)
""")
scores = fresh_db["scores"]
scores.insert({"id": 1, "score": 50})
scores.transform()
assert scores.checks == [
Check("score BETWEEN 0 AND 100", name="valid_score", column="score"),
Check("id != 0", name="nonzero_id"),
]
with pytest.raises(sqlite3.IntegrityError, match="CHECK constraint failed"):
scores.insert({"id": 2, "score": 101})
def test_transform_preserves_check_ending_in_line_comment(fresh_db):
fresh_db.execute("""
CREATE TABLE inventory (
quantity INTEGER,
CHECK (
quantity >= 0 -- Quantity cannot be negative
)
)
""")
inventory = fresh_db["inventory"]
inventory.transform(types={"quantity": float})
assert inventory.checks == [Check("quantity >= 0 -- Quantity cannot be negative")]
with pytest.raises(sqlite3.IntegrityError, match="CHECK constraint failed"):
inventory.insert({"quantity": -1})
def test_transform_renames_columns_inside_check_constraints(fresh_db):
fresh_db.execute("""
CREATE TABLE inventory (
quantity INTEGER CONSTRAINT positive
CHECK(quantity > 0 AND 'quantity' != ''),
maximum INTEGER,
CONSTRAINT within_maximum CHECK(quantity <= maximum)
)
""")
inventory = fresh_db["inventory"]
inventory.insert({"quantity": 2, "maximum": 3})
inventory.transform(rename={"quantity": "amount"})
assert inventory.checks == [
Check(
"amount > 0 AND 'quantity' != ''",
name="positive",
column="amount",
),
Check("amount <= maximum", name="within_maximum"),
]
with pytest.raises(sqlite3.IntegrityError, match="CHECK constraint failed"):
inventory.insert({"amount": 4, "maximum": 3})
def test_transform_check_rewrite_preserves_functions_and_quotes(fresh_db):
fresh_db.execute("""
CREATE TABLE items (
length TEXT,
"old name" TEXT,
CHECK(length("old name") > 0 AND length != '')
)
""")
items = fresh_db["items"]
items.insert({"length": "label", "old name": "hello"})
items.transform(rename={"length": "description", "old name": "new name"})
assert items.checks == [Check("length(\"new name\") > 0 AND description != ''")]
def test_transform_check_rewrite_quotes_keyword_column(fresh_db):
fresh_db.execute("CREATE TABLE t(old_name TEXT CHECK(old_name != ''))")
fresh_db["t"].insert({"old_name": "value"})
fresh_db["t"].transform(rename={"old_name": "select"})
assert fresh_db["t"].checks == [Check("\"select\" != ''", column="select")]
def test_transform_check_rewrite_does_not_rename_collations_or_cast_types(fresh_db):
fresh_db.execute("""
CREATE TABLE t (
nocase TEXT,
kind TEXT,
other TEXT,
CHECK(
other COLLATE nocase != ''
AND CAST(other AS kind) != ''
AND nocase != ''
AND kind != ''
)
)
""")
fresh_db["t"].insert({"nocase": "n", "kind": "k", "other": "o"})
fresh_db["t"].transform(rename={"nocase": "label", "kind": "category"})
check = fresh_db["t"].checks[0].check
assert "COLLATE nocase" in check
assert "AS kind" in check
assert "AND label != ''" in check
assert "AND category != ''" in check
def test_transform_drops_check_owned_by_dropped_column(fresh_db):
fresh_db.execute("""
CREATE TABLE t (
id INTEGER,
obsolete INTEGER CHECK(obsolete > 0),
CHECK(id > 0)
)
""")
fresh_db["t"].insert({"id": 1, "obsolete": 2})
fresh_db["t"].transform(drop={"obsolete"})
assert fresh_db["t"].checks == [Check("id > 0")]
def test_transform_refuses_to_drop_column_used_by_remaining_check(fresh_db):
fresh_db.execute("""
CREATE TABLE ranges (
minimum INTEGER,
maximum INTEGER,
CHECK(minimum <= maximum)
)
""")
ranges = fresh_db["ranges"]
ranges.insert({"minimum": 1, "maximum": 2})
schema_before = ranges.schema
with pytest.raises(
TransformError,
match="Cannot drop column 'maximum'.*CHECK constraint",
):
ranges.transform(drop={"maximum"})
assert ranges.schema == schema_before