Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 52 additions & 1 deletion sqlite_utils/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -4460,8 +4460,59 @@ def build_insert_queries_and_params(
# All columns are in the PK – nothing to update.
do_clause = "DO NOTHING"

# Detect omitted columns with NOT NULL constraints and no defaults.
# SQLite validates NOT NULL constraints during the INSERT evaluation
# before ON CONFLICT triggers. Supplying a correlated scalar subquery
# pulls the existing value so the constraint passes on existing rows,
# while still failing with IntegrityError if the row does not exist.
omitted_not_null_cols: list[str] = []
if self.exists():
lower_all_columns = {c.lower() for c in all_columns}
for col in self.columns:
if col.notnull and col.default_value is None and not col.is_pk:
if col.name.lower() not in lower_all_columns:
omitted_not_null_cols.append(col.name)
if not_null:
lower_all_columns = {c.lower() for c in all_columns}
lower_omitted = {c.lower() for c in omitted_not_null_cols}
for nn in not_null:
if (
nn.lower() not in lower_all_columns
and nn.lower() not in lower_omitted
):
omitted_not_null_cols.append(nn)

if omitted_not_null_cols:
all_insert_cols = list(all_columns) + omitted_not_null_cols
insert_columns_sql = ", ".join(
quote_identifier(c) for c in all_insert_cols
)
where_pk_sql = " AND ".join(
f"{quote_identifier(pk)} = ?" for pk in pk_cols
)
subqueries = [
f"(SELECT {quote_identifier(c)} FROM {quote_identifier(self.name)} WHERE {where_pk_sql})"
for c in omitted_not_null_cols
]
row_placeholder_parts = [
conversions.get(c, "?") for c in all_columns
] + subqueries
single_row_placeholder = f"({', '.join(row_placeholder_parts)})"
row_placeholders_sql = ", ".join(single_row_placeholder for _ in values)

pk_indexes = [all_columns.index(c) for c in pk_cols]
all_flat_params: list[Any] = []
for record_values in values:
all_flat_params.extend(record_values)
row_pk_params = [record_values[i] for i in pk_indexes]
for _ in omitted_not_null_cols:
all_flat_params.extend(row_pk_params)
flat_params = all_flat_params
else:
insert_columns_sql = columns_sql

sql = (
f"INSERT INTO {quote_identifier(self.name)} ({columns_sql}) "
f"INSERT INTO {quote_identifier(self.name)} ({insert_columns_sql}) "
f"VALUES {row_placeholders_sql} "
f"ON CONFLICT({conflict_sql}) {do_clause}"
)
Expand Down
87 changes: 87 additions & 0 deletions tests/test_upsert.py
Original file line number Diff line number Diff line change
Expand Up @@ -185,3 +185,90 @@ def test_upsert_compound_primary_key(fresh_db):
# .upsert_all() with a single item should set .last_pk
table.upsert_all([{"species": "cat", "id": 1, "age": 5}], pk=("species", "id"))
assert ("cat", 1) == table.last_pk


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_omitted_not_null_column_on_existing_row(use_old_upsert):
# https://github.com/simonw/sqlite-utils/issues/878
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("dogs")
table.insert(
{"id": 1, "name": "Cleo", "color": "brown"},
pk="id",
not_null={"name"},
)

# Upserting with omitted NOT NULL column on an existing row must succeed
table.upsert({"id": 1, "color": "black"})
assert table.get(1) == {"id": 1, "name": "Cleo", "color": "black"}


@pytest.mark.parametrize("use_old_upsert", (False, True))
@pytest.mark.parametrize("batch_size", (1, 2))
def test_upsert_all_omitted_not_null_compound_pk(use_old_upsert, batch_size):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("pets")
table.insert_all(
[
{"species": "dog", "id": 1, "name": "Cleo", "breed": "Lab", "age": 4},
{"species": "cat", "id": 1, "name": "Nixie", "breed": "Siamese", "age": 5},
],
pk=("species", "id"),
not_null={"name", "breed"},
)

table.upsert_all(
[
{"species": "dog", "id": 1, "age": 5},
{"species": "cat", "id": 1, "age": 6},
],
batch_size=batch_size,
)

assert table.get(("dog", 1)) == {
"species": "dog",
"id": 1,
"name": "Cleo",
"breed": "Lab",
"age": 5,
}
assert table.get(("cat", 1)) == {
"species": "cat",
"id": 1,
"name": "Nixie",
"breed": "Siamese",
"age": 6,
}


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_omitted_not_null_with_conversions_and_explicit_not_null(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("dogs")
table.insert(
{"id": 1, "name": "Cleo", "color": "brown"},
pk="id",
)
# Upsert with explicit not_null argument and conversion
table.upsert(
{"id": 1, "color": "black"},
not_null=["name"],
conversions={"color": "upper(?)"},
)
assert table.get(1) == {"id": 1, "name": "Cleo", "color": "BLACK"}


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_omitted_not_null_with_check_constraint(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
db.execute("""
CREATE TABLE dogs (
id INTEGER PRIMARY KEY,
name TEXT NOT NULL CHECK(length(name) >= 3),
color TEXT
)
""")
table = db.table("dogs")
table.insert({"id": 1, "name": "Cleo", "color": "brown"})
table.upsert({"id": 1, "color": "black"})
assert table.get(1) == {"id": 1, "name": "Cleo", "color": "black"}
Loading