Refactor .update() to use .get()

.pks introspection now returns [rowid] for rowid tables.
This commit is contained in:
Simon Willison 2019-07-28 15:44:33 +03:00
commit e4a11b1815
2 changed files with 20 additions and 30 deletions

View file

@ -456,31 +456,24 @@ class Table:
@property @property
def pks(self): def pks(self):
return [column.name for column in self.columns if column.is_pk] names = [column.name for column in self.columns if column.is_pk]
if not names:
names = ["rowid"]
return names
def get(self, pk_values): def get(self, pk_values):
if not isinstance(pk_values, (list, tuple)): if not isinstance(pk_values, (list, tuple)):
pk_values = [pk_values] pk_values = [pk_values]
pks = self.pks pks = self.pks
pk_names = [] last_pk = pk_values[0] if len(pks) == 1 else pk_values
if len(pks) == 0: if len(pks) != len(pk_values):
# rowid table
pk_names = ["rowid"]
last_pk = pk_values[0]
elif len(pks) == 1:
pk_names = [pks[0]]
last_pk = pk_values[0]
elif len(pks) > 1:
pk_names = pks
last_pk = pk_values
if len(pk_names) != len(pk_values):
raise NotFoundError( raise NotFoundError(
"Need {} primary key value{}".format( "Need {} primary key value{}".format(
len(pk_names), "" if len(pk_names) == 1 else "s" len(pks), "" if len(pks) == 1 else "s"
) )
) )
wheres = ["[{}] = ?".format(pk_name) for pk_name in pk_names] wheres = ["[{}] = ?".format(pk_name) for pk_name in pks]
rows = self.rows_where(" and ".join(wheres), pk_values) rows = self.rows_where(" and ".join(wheres), pk_values)
try: try:
row = list(rows)[0] row = list(rows)[0]
@ -788,26 +781,15 @@ class Table:
updates = updates or {} updates = updates or {}
if not isinstance(pk_values, (list, tuple)): if not isinstance(pk_values, (list, tuple)):
pk_values = [pk_values] pk_values = [pk_values]
pks = self.pks # Sanity check that the record exists (raises error if not):
pk_names = [] self.get(pk_values)
if len(pks) == 0:
# rowid table
pk_names = ["rowid"]
last_pk = pk_values[0]
elif len(pks) == 1:
pk_names = [pks[0]]
last_pk = pk_values[0]
elif len(pks) > 1:
pk_names = pks
last_pk = pk_values
assert len(pk_names) == len(pk_values)
args = [] args = []
sets = [] sets = []
wheres = [] wheres = []
for key, value in updates.items(): for key, value in updates.items():
sets.append("[{}] = ?".format(key)) sets.append("[{}] = ?".format(key))
args.append(value) args.append(value)
wheres = ["[{}] = ?".format(pk_name) for pk_name in pk_names] wheres = ["[{}] = ?".format(pk_name) for pk_name in self.pks]
args.extend(pk_values) args.extend(pk_values)
sql = "update [{table}] set {sets} where {wheres}".format( sql = "update [{table}] set {sets} where {wheres}".format(
table=self.name, sets=", ".join(sets), wheres=" and ".join(wheres) table=self.name, sets=", ".join(sets), wheres=" and ".join(wheres)
@ -816,7 +798,7 @@ class Table:
rowcount = self.db.conn.execute(sql, args).rowcount rowcount = self.db.conn.execute(sql, args).rowcount
# TODO: Test this works (rolls back) - use better exception: # TODO: Test this works (rolls back) - use better exception:
assert rowcount == 1 assert rowcount == 1
self.last_pk = last_pk self.last_pk = pk_values[0] if len(self.pks) == 1 else pk_values
return self return self
def insert( def insert(

View file

@ -98,3 +98,11 @@ def test_guess_foreign_table(fresh_db, column, expected_table_guess):
fresh_db.create_table("authors", {"name": str}) fresh_db.create_table("authors", {"name": str})
fresh_db.create_table("genre", {"name": str}) fresh_db.create_table("genre", {"name": str})
assert expected_table_guess == fresh_db["books"].guess_foreign_table(column) assert expected_table_guess == fresh_db["books"].guess_foreign_table(column)
@pytest.mark.parametrize(
"pk,expected", ((None, ["rowid"]), ("id", ["id"]), (["id", "id2"], ["id", "id2"]))
)
def test_pks(fresh_db, pk, expected):
fresh_db["foo"].insert_all([{"id": 1, "id2": 2}], pk=pk)
assert expected == fresh_db["foo"].pks