diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 9b179ac..d7fe44a 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -78,6 +78,9 @@ class NoObviousTable(Exception): class BadPrimaryKey(Exception): pass +class NotFoundError(Exception): + pass + class Database: def __init__(self, filename_or_conn): @@ -382,6 +385,30 @@ class Table: def pks(self): return [column.name for column in self.columns if column.is_pk] + def get(self, pk_values): + if not isinstance(pk_values, (list, tuple)): + pk_values = [pk_values] + pks = self.pks + pk_names = [] + 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 + wheres = ["[{}] = ?".format(pk_name) for pk_name in pk_names] + rows = self.rows_where(' and '.join(wheres), pk_values) + try: + row = list(rows)[0] + self.last_pk = last_pk + return row + except IndexError: + raise NotFoundError + @property def foreign_keys(self): fks = [] diff --git a/tests/test_get.py b/tests/test_get.py index 75ddaef..bd79490 100644 --- a/tests/test_get.py +++ b/tests/test_get.py @@ -1,6 +1,20 @@ import pytest +def test_get_rowid(fresh_db): + dogs = fresh_db["dogs"] + cleo = {"name": "Cleo", "age": 4} + row_id = dogs.insert(cleo).last_rowid + assert cleo == dogs.get(row_id) + + +def test_get_primary_key(fresh_db): + dogs = fresh_db["dogs"] + cleo = {"name": "Cleo", "age": 4, "id": 5} + row_id = dogs.insert(cleo, pk="id").last_pk + assert cleo == dogs.get(5) + + @pytest.mark.parametrize( "where,where_args,expected_ids", [