mirror of
https://github.com/simonw/sqlite-utils.git
synced 2026-09-28 12:54:15 +02:00
Refactor .update() to use .get()
.pks introspection now returns [rowid] for rowid tables.
This commit is contained in:
parent
455071f3c5
commit
e4a11b1815
2 changed files with 20 additions and 30 deletions
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue