diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 6ba718a..2430f0d 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -897,7 +897,11 @@ class Table(Queryable): first_column = columns[0] pks = self.pks lookup_table = self.db[table] - for row in self.rows: + if pks == ["rowid"]: + rows_iter = self.rows_where(select="rowid, *") + else: + rows_iter = self.rows + for row in rows_iter: row_pks = tuple(row[pk] for pk in pks) lookups = {rename.get(column) or column: row[column] for column in columns} self.update(row_pks, {first_column: lookup_table.lookup(lookups)}) diff --git a/tests/test_extract.py b/tests/test_extract.py index 4855441..4108a5a 100644 --- a/tests/test_extract.py +++ b/tests/test_extract.py @@ -61,9 +61,9 @@ def test_extract_multiple_columns_with_rename(fresh_db): pk="id", ) - fresh_db["tree"].extract(["common_name", "latin_name"], rename={ - "common_name": "name" - }) + fresh_db["tree"].extract( + ["common_name", "latin_name"], rename={"common_name": "name"} + ) assert fresh_db["tree"].schema == ( 'CREATE TABLE "tree" (\n' " [id] INTEGER PRIMARY KEY,\n" @@ -105,3 +105,22 @@ def test_extract_invalid_columns(fresh_db): ) with pytest.raises(InvalidColumns): fresh_db["tree"].extract(["bad_column"]) + + +def test_extract_rowid_table(fresh_db): + fresh_db["tree"].insert( + { + "name": "Tree 1", + "common_name": "Palm", + "latin_name": "Arecaceae", + } + ) + fresh_db["tree"].extract(["common_name", "latin_name"]) + assert fresh_db["tree"].schema == ( + 'CREATE TABLE "tree" (\n' + " [rowid] INTEGER PRIMARY KEY,\n" + " [name] TEXT,\n" + " [common_name_latin_name_id] INTEGER,\n" + " FOREIGN KEY(common_name_latin_name_id) REFERENCES common_name_latin_name(id)\n" + ")" + ) diff --git a/tests/test_rows.py b/tests/test_rows.py index 8885802..603f254 100644 --- a/tests/test_rows.py +++ b/tests/test_rows.py @@ -26,7 +26,9 @@ def test_rows_where(where, where_args, expected_ids, fresh_db): ], pk="id", ) - assert expected_ids == {r["id"] for r in table.rows_where(where, where_args, select="id")} + assert expected_ids == { + r["id"] for r in table.rows_where(where, where_args, select="id") + } @pytest.mark.parametrize(