diff --git a/docs/python-api.rst b/docs/python-api.rst index def0b14..c0eb593 100644 --- a/docs/python-api.rst +++ b/docs/python-api.rst @@ -231,6 +231,40 @@ This method also accepts ``offset=`` and ``limit=`` arguments, for specifying an ... print(row) {'id': 1, 'age': 4, 'name': 'Cleo'} +.. _python_api_pks_and_rows_where: + +Listing rows with their primary keys +==================================== + +Sometimes it can be useful to retrieve the primary key along with each row, in order to pass that key (or primary key tuple) to the ``.get()`` or ``.update()`` methods. + +The ``.pks_and_rows_where()`` method takes the same signature as ``.rows_where()`` (with the exception of the ``select=`` parameter) but returns a generator that yields pairs of ``(primary key, row dictionary)``. + +The primary key value will usually be a single value but can also be a tuple if the table has a compound primary key. + +If the table is a ``rowid`` table (with no explicit primary key column) then that ID will be returned. + +:: + + >>> db = sqlite_utils.Database(memory=True) + >>> db["dogs"].insert({"name": "Cleo"}) + >>> for pk, row in db["dogs"].pks_and_rows_where(): + ... print(pk, row) + 1 {'rowid': 1, 'name': 'Cleo'} + + >>> db["dogs_with_pk"].insert({"id": 5, "name": "Cleo"}, pk="id") + >>> for pk, row in db["dogs_with_pk"].pks_and_rows_where(): + ... print(pk, row) + 5 {'id': 5, 'name': 'Cleo'} + + >>> db["dogs_with_compound_pk"].insert( + ... {"species": "dog", "id": 3, "name": "Cleo"}, + ... pk=("species", "id") + ... ) + >>> for pk, row in db["dogs_with_compound_pk"].pks_and_rows_where(): + ... print(pk, row) + ('dog', 3) {'species': 'dog', 'id': 3, 'name': 'Cleo'} + .. _python_api_get: Retrieving a specific record diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 971bf6f..97c55b7 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -677,6 +677,34 @@ class Queryable: for row in cursor: yield dict(zip(columns, row)) + def pks_and_rows_where( + self, + where=None, + where_args=None, + order_by=None, + limit=None, + offset=None, + ): + "Like .rows_where() but returns (pk, row) pairs - pk can be a single value or tuple" + column_names = [column.name for column in self.columns] + pks = [column.name for column in self.columns if column.is_pk] + if not pks: + column_names.insert(0, "rowid") + pks = ["rowid"] + select = ",".join("[{}]".format(column_name) for column_name in column_names) + for row in self.rows_where( + select=select, + where=where, + where_args=where_args, + order_by=order_by, + limit=limit, + offset=offset, + ): + row_pk = tuple(row[pk] for pk in pks) + if len(row_pk) == 1: + row_pk = row_pk[0] + yield row_pk, row + @property def columns(self): if not self.exists(): diff --git a/tests/test_rows.py b/tests/test_rows.py index 6d995e2..3d33ecb 100644 --- a/tests/test_rows.py +++ b/tests/test_rows.py @@ -68,3 +68,40 @@ def test_rows_where_offset_limit(fresh_db, offset, limit, expected): assert expected == [ r["id"] for r in table.rows_where(offset=offset, limit=limit, order_by="id") ] + + +def test_pks_and_rows_where_rowid(fresh_db): + table = fresh_db["rowid_table"] + table.insert_all({"number": i + 10} for i in range(3)) + pks_and_rows = list(table.pks_and_rows_where()) + assert pks_and_rows == [ + (1, {"rowid": 1, "number": 10}), + (2, {"rowid": 2, "number": 11}), + (3, {"rowid": 3, "number": 12}), + ] + + +def test_pks_and_rows_where_simple_pk(fresh_db): + table = fresh_db["simple_pk_table"] + table.insert_all(({"id": i + 10} for i in range(3)), pk="id") + pks_and_rows = list(table.pks_and_rows_where()) + assert pks_and_rows == [ + (10, {"id": 10}), + (11, {"id": 11}), + (12, {"id": 12}), + ] + + +def test_pks_and_rows_where_compound_pk(fresh_db): + table = fresh_db["compound_pk_table"] + table.insert_all( + ({"type": "number", "number": i, "plusone": i + 1} for i in range(3)), + pk=("type", "number"), + ) + pks_and_rows = list(table.pks_and_rows_where()) + assert pks_and_rows == [ + (("number", 0), {"type": "number", "number": 0, "plusone": 1}), + (("number", 1), {"type": "number", "number": 1, "plusone": 2}), + (("number", 2), {"type": "number", "number": 2, "plusone": 3}), + ] + assert False