Use ON CONFLICT for upsert, refs #652

* New upsert implementation, refs #652
* supports_strict now caches on self._supports_strict

PR: https://github.com/simonw/sqlite-utils/pull/653
This commit is contained in:
Simon Willison 2025-05-08 20:37:49 -07:00 committed by GitHub
commit 8e7d018fa2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 209 additions and 110 deletions

View file

@ -304,6 +304,8 @@ class Database:
``sql, parameters`` every time a SQL query is executed
:param use_counts_table: set to ``True`` to use a cached counts table, if available. See
:ref:`python_api_cached_table_counts`
:param use_old_upsert: set to ``True`` to force the older upsert implementation. See
:ref:`python_api_old_upsert`
:param strict: Apply STRICT mode to all created tables (unless overridden)
"""
@ -320,10 +322,12 @@ class Database:
tracer: Optional[Callable] = None,
use_counts_table: bool = False,
execute_plugins: bool = True,
use_old_upsert: bool = False,
strict: bool = False,
):
self.memory_name = None
self.memory = False
self.use_old_upsert = use_old_upsert
assert (filename_or_conn is not None and (not memory and not memory_name)) or (
filename_or_conn is None and (memory or memory_name)
), "Either specify a filename_or_conn or pass memory=True"
@ -671,16 +675,46 @@ class Database:
@property
def supports_strict(self) -> bool:
"Does this database support STRICT mode?"
try:
if not hasattr(self, "_supports_strict"):
try:
table_name = "t{}".format(secrets.token_hex(16))
with self.conn:
self.conn.execute(
"create table {} (name text) strict".format(table_name)
)
self.conn.execute("drop table {}".format(table_name))
self._supports_strict = True
except Exception:
self._supports_strict = False
return self._supports_strict
@property
def supports_on_conflict(self) -> bool:
# SQLite's upsert is implemented as INSERT INTO ... ON CONFLICT DO ...
if not hasattr(self, "_supports_on_conflict"):
table_name = "t{}".format(secrets.token_hex(16))
with self.conn:
self.conn.execute(
"create table {} (name text) strict".format(table_name)
)
self.conn.execute("drop table {}".format(table_name))
return True
except Exception:
return False
try:
with self.conn:
self.conn.execute(
"create table {} (id integer primary key, name text)".format(
table_name
)
)
self.conn.execute(
"insert into {} (id, name) values (1, 'one')".format(table_name)
)
self.conn.execute(
(
"insert into {} (id, name) values (1, 'two') "
"on conflict do update set name = 'two'"
).format(table_name)
)
self._supports_on_conflict = True
except Exception:
self._supports_on_conflict = False
finally:
self.conn.execute("drop table if exists {}".format(table_name))
return self._supports_on_conflict
@property
def sqlite_version(self) -> Tuple[int, ...]:
@ -2966,13 +3000,19 @@ class Table(Queryable):
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 = []
"""
Given a list ``chunk`` of records that should be written to *this* table,
return a list of ``(sql, parameters)`` 2-tuples which, when executed in
order, perform the desired INSERT / UPSERT / REPLACE operation.
"""
if hash_id_columns and hash_id is None:
hash_id = "id"
extracts = resolve_extracts(extracts)
# Build a row-list ready for executemany-style flattening
values = []
for record in chunk:
record_values = []
for key in all_columns:
@ -2992,76 +3032,103 @@ class Table(Queryable):
record_values.append(value)
values.append(record_values)
queries_and_params = []
if upsert:
if isinstance(pk, str):
pks = [pk]
columns_sql = ", ".join(f"[{c}]" for c in all_columns)
placeholder_expr = ", ".join(conversions.get(c, "?") for c in all_columns)
row_placeholders_sql = ", ".join(f"({placeholder_expr})" for _ in values)
flat_params = list(itertools.chain.from_iterable(values))
# replace=True mean INSERT OR REPLACE INTO
if replace:
sql = (
f"INSERT OR REPLACE INTO [{self.name}] "
f"({columns_sql}) VALUES {row_placeholders_sql}"
)
return [(sql, flat_params)]
# If not an upsert it's an INSERT, maybe with OR IGNORE
if not upsert:
or_ignore = ""
if ignore:
or_ignore = " OR IGNORE"
sql = (
f"INSERT{or_ignore} INTO [{self.name}] "
f"({columns_sql}) VALUES {row_placeholders_sql}"
)
return [(sql, flat_params)]
# Everything from here on is for upsert=True
pk_cols = [pk] if isinstance(pk, str) else list(pk)
non_pk_cols = [c for c in all_columns if c not in pk_cols]
conflict_sql = ", ".join(f"[{c}]" for c in pk_cols)
if self.db.supports_on_conflict and not self.db.use_old_upsert:
if non_pk_cols:
# DO UPDATE
assignments = []
for c in non_pk_cols:
if c in conversions:
assignments.append(
f"[{c}] = {conversions[c].replace('?', f'excluded.[{c}]')}"
)
else:
assignments.append(f"[{c}] = excluded.[{c}]")
do_clause = "DO UPDATE SET " + ", ".join(assignments)
else:
pks = pk
self.last_pk = None
for record_values in values:
record = dict(zip(all_columns, record_values))
placeholders = list(pks)
# Need to populate not-null columns too, or INSERT OR IGNORE ignores
# them since it ignores the resulting integrity errors
if not_null:
placeholders.extend(not_null)
sql = "INSERT OR IGNORE INTO [{table}]({cols}) VALUES({placeholders});".format(
# All columns are in the PK nothing to update.
do_clause = "DO NOTHING"
sql = (
f"INSERT INTO [{self.name}] ({columns_sql}) "
f"VALUES {row_placeholders_sql} "
f"ON CONFLICT({conflict_sql}) {do_clause}"
)
return [(sql, flat_params)]
# At this point we need compatibility UPSERT for SQLite < 3.24.0
# (INSERT OR IGNORE + second UPDATE stage)
queries_and_params = []
if isinstance(pk, str):
pks = [pk]
else:
pks = pk
self.last_pk = None
for record_values in values:
record = dict(zip(all_columns, record_values))
placeholders = list(pks)
# Need to populate not-null columns too, or INSERT OR IGNORE ignores
# them since it ignores the resulting integrity errors
if not_null:
placeholders.extend(not_null)
sql = "INSERT OR IGNORE INTO [{table}]({cols}) VALUES({placeholders});".format(
table=self.name,
cols=", ".join(["[{}]".format(p) for p in placeholders]),
placeholders=", ".join(["?" for p in placeholders]),
)
queries_and_params.append(
(sql, [record[col] for col in pks] + ["" for _ in (not_null or [])])
)
# UPDATE [book] SET [name] = 'Programming' WHERE [id] = 1001;
set_cols = [col for col in all_columns if col not in pks]
if set_cols:
sql2 = "UPDATE [{table}] SET {pairs} WHERE {wheres}".format(
table=self.name,
cols=", ".join(["[{}]".format(p) for p in placeholders]),
placeholders=", ".join(["?" for p in placeholders]),
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(
(sql, [record[col] for col in pks] + ["" for _ in (not_null or [])])
(
sql2,
[record[col] for col in set_cols] + [record[pk] for pk in pks],
)
)
# UPDATE [book] SET [name] = 'Programming' WHERE [id] = 1001;
set_cols = [col for col in all_columns if col not in pks]
if set_cols:
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)]
# 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]
return queries_and_params
def insert_chunk(
@ -3079,7 +3146,7 @@ class Table(Queryable):
num_records_processed,
replace,
ignore,
):
) -> Optional[sqlite3.Cursor]:
queries_and_params = self.build_insert_queries_and_params(
extracts,
chunk,
@ -3094,9 +3161,8 @@ class Table(Queryable):
replace,
ignore,
)
result = None
with self.db.conn:
result = None
for query, params in queries_and_params:
try:
result = self.db.execute(query, params)
@ -3125,7 +3191,7 @@ class Table(Queryable):
ignore,
)
self.insert_chunk(
result = self.insert_chunk(
alter,
extracts,
second_half,
@ -3143,20 +3209,7 @@ class Table(Queryable):
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
return result
def insert(
self,
@ -3276,6 +3329,7 @@ class Table(Queryable):
if upsert and (not pk and not hash_id):
raise PrimaryKeyRequired("upsert() requires a pk")
assert not (hash_id and pk), "Use either pk= or hash_id="
if hash_id_columns and (hash_id is None):
hash_id = "id"
@ -3307,6 +3361,7 @@ class Table(Queryable):
self.last_pk = None
if truncate and self.exists():
self.db.execute("DELETE FROM [{}];".format(self.name))
result = None
for chunk in chunks(itertools.chain([first_record], records), batch_size):
chunk = list(chunk)
num_records_processed += len(chunk)
@ -3314,6 +3369,12 @@ class Table(Queryable):
if not self.exists():
# Use the first batch to derive the table names
column_types = suggest_column_types(chunk)
if extracts:
for col in extracts:
if col in column_types:
column_types[col] = (
int # This will be an integer foreign key
)
column_types.update(columns or {})
self.create(
column_types,
@ -3341,7 +3402,7 @@ class Table(Queryable):
first = False
self.insert_chunk(
result = self.insert_chunk(
alter,
extracts,
chunk,
@ -3357,6 +3418,33 @@ class Table(Queryable):
ignore,
)
# If we only handled a single row populate self.last_pk
if num_records_processed == 1:
# For an insert we need to use result.lastrowid
if not upsert and result is not None:
self.last_rowid = result.lastrowid
if (hash_id or pk) and self.last_rowid:
# Set self.last_pk to the pk(s) for that 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)
else:
self.last_pk = self.last_rowid
else:
# For an upsert use first_record from earlier
if hash_id:
self.last_pk = hash_record(first_record, hash_id_columns)
else:
self.last_pk = (
first_record[pk]
if isinstance(pk, str)
else tuple(first_record[p] for p in pk)
)
if analyze:
self.analyze()