From 2cfc29931d299cc9527586788c3663f092d4d416 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Tue, 9 Jul 2019 07:36:37 -0700 Subject: [PATCH] WIP --- sqlite_utils/db.py | 49 ++++++++++++++++++++++++++++++++++++++++++---- 1 file changed, 45 insertions(+), 4 deletions(-) diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 3a44cb4..9354e49 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -79,6 +79,10 @@ class BadPrimaryKey(Exception): pass +class RowNotFound(Exception): + pass + + class Database: def __init__(self, filename_or_conn): if isinstance(filename_or_conn, str): @@ -780,12 +784,26 @@ class Table: self.last_pk = None # self.last_rowid will be 0 if a "INSERT OR IGNORE" happened if (hash_id or pk) and self.last_rowid: - self.last_pk = self.db.conn.execute( - "select [{}] from [{}] where rowid = ?".format( - hash_id or pk, self.name + compound_pk = False + if hash_id: + pk_cols = [hash_id] + elif isinstance(pk, (list, tuple)): + pk_cols = pk + compound_pk = True + else: + pk_cols = [pk] + result = self.db.conn.execute( + "select {} from [{}] where rowid = ?".format( + ", ".join("[{}]".format(pk_col) for pk_col in pk_cols), + self.name, ), (self.last_rowid,), - ).fetchone()[0] + ).fetchone() + if compound_pk: + self.last_pk = list(result) + else: + self.last_pk = result[0] + return self def upsert( @@ -836,6 +854,29 @@ class Table: upsert=True, ) + def update(self, pk, updates=None): + updates = updates or {} + row = self.get(pk) + + def get(self, pk): + "pk is the primary key value, or a tuple/list for compound pks" + if not isinstance(pk, (list, tuple)): + pk = [pk] + # Confirm pk has the correct number of values + if len(pk) != len(self.pks): + raise RowNotFound( + "Primary key value needs to match ({})".format( + ", ".join(map(str, self.pks)) + ) + ) + sql = "SELECT * FROM {} WHERE {}".format( + self.name, " AND ".join("[{}] = ?".format(pk) for pk in self.pks) + ) + rows = self.db.execute_returning_dicts(sql, pk) + if len(rows) != 1: + raise RowNotFound("Row not found for {}".format(repr(pk))) + return rows[0] + def add_missing_columns(self, records): needed_columns = self.detect_column_types(records) current_columns = self.columns_dict