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: 37 additions & 16 deletions sqlite_utils/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -4887,22 +4887,43 @@ def insert_all(

first = False

result = self.insert_chunk(
alter,
extracts,
chunk,
all_columns,
hash_id,
hash_id_columns,
upsert,
pk,
not_null,
conversions,
num_records_processed,
replace,
ignore,
list_mode,
)
if upsert and not replace and not list_mode:
# Only adjacent records with the same fields can share an
# upsert statement. Missing fields must not become NULL updates,
# and grouping non-adjacent records would change update order.
insert_groups = [
(
list(group),
[c for c in all_columns if c in keys or c == hash_id],
)
for keys, group in itertools.groupby(chunk, key=frozenset)
]
else:
insert_groups = [(chunk, all_columns)]

split_upsert = len(insert_groups) > 1
with self.db.atomic() if split_upsert else contextlib.nullcontext():
if split_upsert and alter:
# Keep type inference over the original batch: a later
# group may require TEXT where the first suggests INTEGER.
self.add_missing_columns(cast(list[dict[str, Any]], chunk))
for group, group_columns in insert_groups:
result = self.insert_chunk(
alter,
extracts,
group,
group_columns,
hash_id,
hash_id_columns,
upsert,
pk,
not_null,
conversions,
num_records_processed,
replace,
ignore,
list_mode,
)

# If we only handled a single row populate self.last_pk
if num_records_processed == 1:
Expand Down
123 changes: 123 additions & 0 deletions tests/test_upsert.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,129 @@
from sqlite_utils.db import PrimaryKeyRequired


@pytest.mark.parametrize("use_old_upsert", (False, True))
@pytest.mark.parametrize("batch_size", (1, 2, 100))
def test_upsert_all_preserves_omitted_fields(use_old_upsert, batch_size):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("dogs")
table.insert_all(
[{"id": 1, "name": "Cleo", "age": 5}, {"id": 2, "name": "Nixie", "age": 6}],
pk="id",
)
table.upsert_all(
iter([{"id": 1, "age": 7}, {"id": 2, "name": "Nixie II"}]),
batch_size=batch_size,
)
assert list(table.rows_where(order_by="id")) == [
{"id": 1, "name": "Cleo", "age": 7},
{"id": 2, "name": "Nixie II", "age": 6},
]
assert table.last_pk is None


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_all_preserves_mixed_field_record_order(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("dogs")
table.insert({"id": 1, "name": "Cleo", "age": 5}, pk="id")
db.executescript("""
create table audit (age integer);
create trigger record_age after update on dogs begin
insert into audit (age) values (new.age);
end;
""")
table.upsert_all(
[{"id": 1, "age": 7}, {"id": 1, "name": "Cleo II"}, {"id": 1, "age": 8}]
)
assert table.get(1) == {"id": 1, "name": "Cleo II", "age": 8}
assert list(db["audit"].rows) == [{"age": 7}, {"age": 7}, {"age": 8}]


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_all_omission_differs_from_explicit_null(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("dogs")
table.create(
{"id": int, "name": str, "age": int}, pk="id", defaults={"name": "Unknown"}
)
table.insert({"id": 1, "name": "Cleo", "age": 5})
table.upsert_all([{"id": 1, "name": None}, {"id": 2, "age": 6}])
assert list(table.rows_where(order_by="id")) == [
{"id": 1, "name": None, "age": 5},
{"id": 2, "name": "Unknown", "age": 6},
]


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_all_mixed_fields_invalid_pk_rolls_back_batch(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("dogs")
original = {"id": 1, "name": "Cleo", "age": 5}
table.insert(original, pk="id")
with pytest.raises(PrimaryKeyRequired):
table.upsert_all([{"id": 1, "age": 7}, {"name": "Invalid"}])
assert list(table.rows) == [original]


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_all_mixed_fields_hash_id(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("dogs")
table.upsert_all(
[{"name": "Cleo", "age": 5, "color": "black"}], hash_id_columns=["name"]
)
original_id = table.last_pk
table.upsert_all(
[{"name": "Cleo", "age": 7}, {"name": "Cleo", "color": "brown"}],
hash_id_columns=["name"],
)
assert list(table.rows) == [
{"id": original_id, "name": "Cleo", "age": 7, "color": "brown"}
]


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_all_mixed_fields_alter_infers_types_from_whole_batch(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("items").create({"id": int}, pk="id")
table.upsert_all(
[{"id": 1, "code": 1}, {"id": 2, "code": "001", "note": "leading zeros"}],
alter=True,
)
assert table.columns_dict["code"] is str
assert list(table.rows_where(order_by="id")) == [
{"id": 1, "code": "1", "note": None},
{"id": 2, "code": "001", "note": "leading zeros"},
]


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert_all_mixed_fields_compound_pk_and_conversions(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
table = db.table("pets")
table.insert_all(
[
{"species": "dog", "id": 1, "name": "Cleo", "age": 5},
{"species": "cat", "id": 1, "name": "Nixie", "age": 6},
],
pk=("species", "id"),
)
table.upsert_all(
[
{"species": "dog", "id": 1, "age": 7},
{"species": "cat", "id": 1, "name": "new"},
],
conversions={"name": "upper(?)"},
)
assert table.get(("dog", 1)) == {
"species": "dog",
"id": 1,
"name": "Cleo",
"age": 7,
}
assert table.get(("cat", 1)) == {"species": "cat", "id": 1, "name": "NEW", "age": 6}


@pytest.mark.parametrize("use_old_upsert", (False, True))
def test_upsert(use_old_upsert):
db = Database(memory=True, use_old_upsert=use_old_upsert)
Expand Down
Loading