diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 0846965..9be09f6 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -287,24 +287,24 @@ class Database: (table, column, other_table, other_column) ) - # Construct list of args to be used with "UPDATE sqlite_master SET sql = ? WHERE name = ?" - sql_args = [] + # Construct SQL for use with "UPDATE sqlite_master SET sql = ? WHERE name = ?" + table_sql = {} for table, column, other_table, other_column in foreign_keys: - old_sql = self[table].schema + old_sql = table_sql.get(table, self[table].schema) extra_sql = ",\n FOREIGN KEY({column}) REFERENCES {other_table}({other_column})\n".format( column=column, other_table=other_table, other_column=other_column ) # Stick that bit in at the very end just before the closing ')' last_paren = old_sql.rindex(")") new_sql = old_sql[:last_paren].strip() + extra_sql + old_sql[last_paren:] - sql_args.append((new_sql, table)) + table_sql[table] = new_sql # And execute it all within a single transaction with self.conn: cursor = self.conn.cursor() schema_version = cursor.execute("PRAGMA schema_version").fetchone()[0] cursor.execute("PRAGMA writable_schema = 1") - for new_sql, table_name in sql_args: + for table_name, new_sql in table_sql.items(): cursor.execute( "UPDATE sqlite_master SET sql = ? WHERE name = ?", (new_sql, table_name), diff --git a/tests/test_create.py b/tests/test_create.py index 04b1fbf..88d7676 100644 --- a/tests/test_create.py +++ b/tests/test_create.py @@ -303,6 +303,34 @@ def test_add_foreign_key_error_if_already_exists(fresh_db): assert "Foreign key already exists for author_id => authors.id" == ex.value.args[0] +def test_add_foreign_keys(fresh_db): + fresh_db["authors"].insert_all( + [{"id": 1, "name": "Sally"}, {"id": 2, "name": "Asheesh"}], pk="id" + ) + fresh_db["categories"].insert_all([{"id": 1, "name": "Wildlife"}], pk="id") + fresh_db["books"].insert_all( + [{"title": "Hedgehogs of the world", "author_id": 1, "category_id": 1}] + ) + assert [] == fresh_db["books"].foreign_keys + fresh_db.add_foreign_keys( + [ + ("books", "author_id", "authors", "id"), + ("books", "category_id", "categories", "id"), + ] + ) + assert [ + ForeignKey( + table="books", column="author_id", other_table="authors", other_column="id" + ), + ForeignKey( + table="books", + column="category_id", + other_table="categories", + other_column="id", + ), + ] == sorted(fresh_db["books"].foreign_keys) + + def test_add_column_foreign_key(fresh_db): fresh_db.create_table("dogs", {"name": str}) fresh_db.create_table("breeds", {"name": str})