diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index f938ef4..86f417b 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -1028,6 +1028,186 @@ class Table(Queryable): self.last_pk = pk_values[0] if len(self.pks) == 1 else pk_values return self + def build_insert_queries_and_params( + self, + extracts, + chunk, + all_columns, + hash_id, + upsert, + pk, + conversions, + num_records_processed, + replace, + ignore, + ): + # values is the list of insert data that is passed to the + # .execute() method - but some of them may be replaced by + # new primary keys if we are extracting any columns. + values = [] + extracts = resolve_extracts(extracts) + for record in chunk: + record_values = [] + for key in all_columns: + value = jsonify_if_needed( + record.get(key, None if key != hash_id else _hash(record)) + ) + if key in extracts: + extract_table = extracts[key] + value = self.db[extract_table].lookup({"value": value}) + record_values.append(value) + values.append(record_values) + + queries_and_params = [] + if upsert: + if isinstance(pk, str): + pks = [pk] + else: + pks = pk + self.last_pk = None + for record_values in values: + # TODO: make more efficient: + record = dict(zip(all_columns, record_values)) + sql = "INSERT OR IGNORE INTO [{table}]({pks}) VALUES({pk_placeholders});".format( + table=self.name, + pks=", ".join(["[{}]".format(p) for p in pks]), + pk_placeholders=", ".join(["?" for p in pks]), + ) + queries_and_params.append((sql, [record[col] for col in pks])) + # UPDATE [book] SET [name] = 'Programming' WHERE [id] = 1001; + set_cols = [col for col in all_columns if col not in pks] + sql2 = "UPDATE [{table}] SET {pairs} WHERE {wheres}".format( + table=self.name, + pairs=", ".join( + "[{}] = {}".format(col, conversions.get(col, "?")) + for col in set_cols + ), + wheres=" AND ".join("[{}] = ?".format(pk) for pk in pks), + ) + queries_and_params.append( + ( + sql2, + [record[col] for col in set_cols] + [record[pk] for pk in pks], + ) + ) + # We can populate .last_pk right here + if num_records_processed == 1: + self.last_pk = tuple(record[pk] for pk in pks) + if len(self.last_pk) == 1: + self.last_pk = self.last_pk[0] + + else: + or_what = "" + if replace: + or_what = "OR REPLACE " + elif ignore: + or_what = "OR IGNORE " + sql = """ + INSERT {or_what}INTO [{table}] ({columns}) VALUES {rows}; + """.strip().format( + or_what=or_what, + table=self.name, + columns=", ".join("[{}]".format(c) for c in all_columns), + rows=", ".join( + "({placeholders})".format( + placeholders=", ".join( + [conversions.get(col, "?") for col in all_columns] + ) + ) + for record in chunk + ), + ) + flat_values = list(itertools.chain(*values)) + queries_and_params = [(sql, flat_values)] + + return queries_and_params + + def insert_chunk( + self, + alter, + extracts, + chunk, + all_columns, + hash_id, + upsert, + pk, + conversions, + num_records_processed, + replace, + ignore, + ): + queries_and_params = self.build_insert_queries_and_params( + extracts, + chunk, + all_columns, + hash_id, + upsert, + pk, + conversions, + num_records_processed, + replace, + ignore, + ) + + with self.db.conn: + for query, params in queries_and_params: + try: + result = self.db.execute(query, params) + except OperationalError as e: + if alter and (" column" in e.args[0]): + # Attempt to add any missing columns, then try again + self.add_missing_columns(chunk) + result = self.db.execute(query, params) + elif e.args[0] == "too many SQL variables": + + first_half = chunk[: len(chunk) // 2] + second_half = chunk[len(chunk) // 2 :] + + self.insert_chunk( + alter, + extracts, + first_half, + all_columns, + hash_id, + upsert, + pk, + conversions, + num_records_processed, + replace, + ignore, + ) + + self.insert_chunk( + alter, + extracts, + second_half, + all_columns, + hash_id, + upsert, + pk, + conversions, + num_records_processed, + replace, + ignore, + ) + + else: + raise + if num_records_processed == 1 and not upsert: + self.last_rowid = result.lastrowid + self.last_pk = self.last_rowid + # self.last_rowid will be 0 if a "INSERT OR IGNORE" happened + if (hash_id or pk) and self.last_rowid: + row = list(self.rows_where("rowid = ?", [self.last_rowid]))[0] + if hash_id: + self.last_pk = row[hash_id] + elif isinstance(pk, str): + self.last_pk = row[pk] + else: + self.last_pk = tuple(row[p] for p in pk) + + return + def insert( self, record, @@ -1161,110 +1341,21 @@ class Table(Queryable): validate_column_names(all_columns) first = False - # values is the list of insert data that is passed to the - # .execute() method - but some of them may be replaced by - # new primary keys if we are extracting any columns. - values = [] - extracts = resolve_extracts(extracts) - for record in chunk: - record_values = [] - for key in all_columns: - value = jsonify_if_needed( - record.get(key, None if key != hash_id else _hash(record)) - ) - if key in extracts: - extract_table = extracts[key] - value = self.db[extract_table].lookup({"value": value}) - record_values.append(value) - values.append(record_values) - queries_and_params = [] - if upsert: - if isinstance(pk, str): - pks = [pk] - else: - pks = pk - self.last_pk = None - for record_values in values: - # TODO: make more efficient: - record = dict(zip(all_columns, record_values)) - params = [] - sql = "INSERT OR IGNORE INTO [{table}]({pks}) VALUES({pk_placeholders});".format( - table=self.name, - pks=", ".join(["[{}]".format(p) for p in pks]), - pk_placeholders=", ".join(["?" for p in pks]), - ) - queries_and_params.append((sql, [record[col] for col in pks])) - # UPDATE [book] SET [name] = 'Programming' WHERE [id] = 1001; - set_cols = [col for col in all_columns if col not in pks] - sql2 = "UPDATE [{table}] SET {pairs} WHERE {wheres}".format( - table=self.name, - pairs=", ".join( - "[{}] = {}".format(col, conversions.get(col, "?")) - for col in set_cols - ), - wheres=" AND ".join("[{}] = ?".format(pk) for pk in pks), - ) - queries_and_params.append( - ( - sql2, - [record[col] for col in set_cols] - + [record[pk] for pk in pks], - ) - ) - # We can populate .last_pk right here - if num_records_processed == 1: - self.last_pk = tuple(record[pk] for pk in pks) - if len(self.last_pk) == 1: - self.last_pk = self.last_pk[0] + self.insert_chunk( + alter, + extracts, + chunk, + all_columns, + hash_id, + upsert, + pk, + conversions, + num_records_processed, + replace, + ignore, + ) - else: - or_what = "" - if replace: - or_what = "OR REPLACE " - elif ignore: - or_what = "OR IGNORE " - sql = """ - INSERT {or_what}INTO [{table}] ({columns}) VALUES {rows}; - """.strip().format( - or_what=or_what, - table=self.name, - columns=", ".join("[{}]".format(c) for c in all_columns), - rows=", ".join( - "({placeholders})".format( - placeholders=", ".join( - [conversions.get(col, "?") for col in all_columns] - ) - ) - for record in chunk - ), - ) - flat_values = list(itertools.chain(*values)) - queries_and_params = [(sql, flat_values)] - - with self.db.conn: - for query, params in queries_and_params: - try: - result = self.db.execute(query, params) - except OperationalError as e: - if alter and (" column" in e.args[0]): - # Attempt to add any missing columns, then try again - self.add_missing_columns(chunk) - result = self.db.execute(query, params) - else: - raise - if num_records_processed == 1 and not upsert: - self.last_rowid = result.lastrowid - self.last_pk = self.last_rowid - # self.last_rowid will be 0 if a "INSERT OR IGNORE" happened - if (hash_id or pk) and self.last_rowid: - row = list(self.rows_where("rowid = ?", [self.last_rowid]))[0] - if hash_id: - self.last_pk = row[hash_id] - elif isinstance(pk, str): - self.last_pk = row[pk] - else: - self.last_pk = tuple(row[p] for p in pk) return self def upsert( diff --git a/tests/test_create.py b/tests/test_create.py index 30a3d0e..e139557 100644 --- a/tests/test_create.py +++ b/tests/test_create.py @@ -584,6 +584,30 @@ def test_error_if_more_than_999_columns(fresh_db, num_columns, should_error): fresh_db["big"].insert(record) +def test_columns_not_in_first_record_should_not_cause_batch_to_be_too_large(fresh_db): + # https://github.com/simonw/sqlite-utils/issues/145 + # sqlite on homebrew and Debian/Ubuntu etc. is typically compiled with + # SQLITE_MAX_VARIABLE_NUMBER set to 250,000, so we need to exceed this value to + # trigger the error on these systems. + THRESHOLD = 250000 + batch_size = 999 + extra_columns = 1 + (THRESHOLD - 1) // (batch_size - 1) + records = [ + {"c0": "first record"}, # one column in first record -> batch size = 999 + # fill out the batch with 99 records with enough columns to exceed THRESHOLD + *[ + dict([("c{}".format(i), j) for i in range(extra_columns)]) + for j in range(batch_size - 1) + ], + ] + try: + fresh_db["too_many_columns"].insert_all( + records, alter=True, batch_size=batch_size + ) + except sqlite3.OperationalError: + raise + + @pytest.mark.parametrize( "columns,index_name,expected_index", (