diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index c011d9bb0..77205028c 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -4353,16 +4353,32 @@ def build_insert_queries_and_params( if list_mode: # In list mode, records are already lists of values num_columns = len(all_columns) + # With hash_id, all_columns has the hash column prepended, but the + # records themselves only carry the user-declared columns. + user_columns = all_columns[1:] if hash_id else all_columns + user_num_columns = len(user_columns) has_extracts = bool(extracts) for record in chunk: # Pad short records with None, truncate long ones record_len = len(record) - if record_len < num_columns: + if record_len < user_num_columns: record_values = [jsonify_if_needed(v) for v in record] + [None] * ( - num_columns - record_len + user_num_columns - record_len ) else: - record_values = [jsonify_if_needed(v) for v in record[:num_columns]] + record_values = [ + jsonify_if_needed(v) for v in record[:user_num_columns] + ] + if hash_id: + # Compute the hash from the user-declared values so the + # hash_id column gets the right key instead of shifting + # every value one column over. + record_values.insert( + 0, + hash_record( + dict(zip(user_columns, record_values)), hash_id_columns + ), + ) # Only process extracts if there are any if has_extracts: for i, key in enumerate(all_columns): diff --git a/tests/test_list_mode.py b/tests/test_list_mode.py index 75f5a7613..7e9646313 100644 --- a/tests/test_list_mode.py +++ b/tests/test_list_mode.py @@ -5,6 +5,7 @@ import pytest from sqlite_utils import Database +from sqlite_utils.utils import hash_record def test_insert_all_list_mode_basic(): @@ -287,3 +288,39 @@ def upsert_data(): # Verify last_pk is populated correctly assert table.last_pk == 1 + + +def test_insert_all_list_mode_with_hash_id_aligns_values(): + """hash_id in list mode must hash user values, not shift every column left.""" + db = Database(memory=True) + table = db.table("items") + table.insert_all([["a", "b"], [1, "x"], [2, "y"]], hash_id="id") + + rows = list(table.rows) + assert '"a" INTEGER' in table.schema + assert '"b" TEXT' in table.schema + assert rows[0] == {"id": hash_record({"a": 1, "b": "x"}), "a": 1, "b": "x"} + assert rows[1] == {"id": hash_record({"a": 2, "b": "y"}), "a": 2, "b": "y"} + + +def test_insert_all_list_mode_with_hash_id_pads_and_truncates(): + """Short rows are padded and long rows truncated before hashing.""" + db = Database(memory=True) + table = db.table("items") + table.insert_all([["c1", "c2"], [5], [6, 7, 999]], hash_id="id") + + rows = list(table.rows) + assert rows[0] == {"id": hash_record({"c1": 5, "c2": None}), "c1": 5, "c2": None} + assert rows[1] == {"id": hash_record({"c1": 6, "c2": 7}), "c1": 6, "c2": 7} + + +def test_upsert_all_list_mode_with_hash_id(): + db = Database(memory=True) + table = db.table("data") + table.upsert_all([["k", "v"], [1, "a"], [1, "b"]], hash_id="id") + + rows = list(table.rows) + assert rows[0]["id"] == hash_record({"k": 1, "v": "a"}) + assert rows[0]["k"] == 1 and rows[0]["v"] == "a" + assert rows[1]["id"] == hash_record({"k": 1, "v": "b"}) + assert rows[1]["k"] == 1 and rows[1]["v"] == "b"