diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index c011d9bb0..c8a54e858 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -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: diff --git a/tests/test_upsert.py b/tests/test_upsert.py index 8274557fa..05b87f7b2 100644 --- a/tests/test_upsert.py +++ b/tests/test_upsert.py @@ -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)