From b75edf4b30dd500747f6f9f5f1dcb52c03a5b442 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sun, 5 Jul 2026 14:37:32 -0700 Subject: [PATCH] add_foreign_key() accepts compound keys, composite FK indexes, refs #594 table.add_foreign_key() and db.add_foreign_keys() now accept lists of column names to create compound foreign keys: table.add_foreign_key( ["campus_name", "dept_code"], "departments" ) Omitting the other columns uses the compound primary key of the other table. Duplicate detection compares the full column lists. db.index_foreign_keys() now creates a single composite index across the columns of a compound foreign key, rather than an index per column. Co-Authored-By: Claude Fable 5 --- docs/python-api.rst | 12 ++++ sqlite_utils/db.py | 138 ++++++++++++++++++++++++------------- tests/test_foreign_keys.py | 74 ++++++++++++++++++++ 3 files changed, 177 insertions(+), 47 deletions(-) diff --git a/docs/python-api.rst b/docs/python-api.rst index 465df79..882f952 100644 --- a/docs/python-api.rst +++ b/docs/python-api.rst @@ -1562,6 +1562,16 @@ To ignore the case where the key already exists, use ``ignore=True``: db.table("books").add_foreign_key("author_id", "authors", "id", ignore=True) +To add a compound foreign key, pass lists of columns: + +.. code-block:: python + + db.table("courses").add_foreign_key( + ["campus_name", "dept_code"], "departments", ["campus_name", "dept_code"] + ) + +As with single columns, omitting the other columns will use the compound primary key of the other table. ``other_table`` must always be specified for a compound foreign key. + .. _python_api_add_foreign_keys: Adding multiple foreign key constraints at once @@ -1591,6 +1601,8 @@ If you want to ensure that every foreign key column in your database has a corre db.index_foreign_keys() +Compound foreign keys get a single composite index across their columns. + .. _python_api_drop: Dropping a table or view diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index c37f750..3c78e3b 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -1590,13 +1590,14 @@ class Database: return candidates def add_foreign_keys( - self, foreign_keys: Iterable[Tuple[str, str, str, str]] + self, foreign_keys: Iterable[Union[ForeignKey, Tuple[str, Any, str, Any]]] ) -> None: """ See :ref:`python_api_add_foreign_keys`. :param foreign_keys: A list of ``(table, column, other_table, other_column)`` - tuples + tuples - for compound foreign keys, ``column`` and ``other_column`` can + be lists of column names """ # foreign_keys is a list of explicit 4-tuples if not all( @@ -1609,44 +1610,64 @@ class Database: "(table, column, other_table, other_column)" ) - foreign_keys_to_create = [] + foreign_keys_to_create: List[Tuple[str, Any, str, Any]] = [] # Verify that all tables and columns exist for fk in foreign_keys: if isinstance(fk, ForeignKey): - table, column, other_table, other_column = ( + table, columns, other_table, other_columns = ( fk.table, - fk.column, + fk.columns, fk.other_table, - fk.other_column, + fk.other_columns, ) else: - table, column, other_table, other_column = fk + table, column_or_columns, other_table, other_column_or_columns = fk + # Compound foreign keys use lists of columns + columns = ( + [column_or_columns] + if isinstance(column_or_columns, str) + else list(column_or_columns) + ) + other_columns = ( + [other_column_or_columns] + if isinstance(other_column_or_columns, str) + else list(other_column_or_columns) + ) if not self.table(table).exists(): raise AlterError("No such table: {}".format(table)) table_obj = self.table(table) - if column not in table_obj.columns_dict: - raise AlterError("No such column: {} in {}".format(column, table)) + for column in columns: + if column not in table_obj.columns_dict: + raise AlterError("No such column: {} in {}".format(column, table)) if not self[other_table].exists(): raise AlterError("No such other_table: {}".format(other_table)) - if ( - other_column != "rowid" - and other_column not in self[other_table].columns_dict - ): - raise AlterError( - "No such other_column: {} in {}".format(other_column, other_table) - ) + for other_column in other_columns: + if ( + other_column != "rowid" + and other_column not in self[other_table].columns_dict + ): + raise AlterError( + "No such other_column: {} in {}".format( + other_column, other_table + ) + ) # We will silently skip foreign keys that exist already if not any( fk for fk in table_obj.foreign_keys - if fk.column == column + if fk.columns == columns and fk.other_table == other_table - and fk.other_column == other_column + and fk.other_columns == other_columns ): - foreign_keys_to_create.append( - (table, column, other_table, other_column) - ) + if len(columns) == 1: + foreign_keys_to_create.append( + (table, columns[0], other_table, other_columns[0]) + ) + else: + foreign_keys_to_create.append( + (table, columns, other_table, other_columns) + ) # Group them by table by_table: Dict[str, List] = {} @@ -1663,15 +1684,12 @@ class Database: "Create indexes for every foreign key column on every table in the database." for table_name in self.table_names(): table = self.table(table_name) - existing_indexes = { - i.columns[0] for i in table.indexes if len(i.columns) == 1 - } + existing_indexes = {tuple(i.columns) for i in table.indexes} for fk in table.foreign_keys: - # Compound foreign keys expose their columns via fk.columns; - # single-column keys yield a one-item list - for column in fk.columns: - if column not in existing_indexes: - table.create_index([column], find_unique_name=True) + # A compound foreign key gets a single composite index + if tuple(fk.columns) not in existing_indexes: + table.create_index(fk.columns, find_unique_name=True) + existing_indexes.add(tuple(fk.columns)) def vacuum(self) -> None: "Run a SQLite ``VACUUM`` against the database." @@ -2862,52 +2880,78 @@ class Table(Queryable): def add_foreign_key( self, - column: str, + column: Union[str, List[str]], other_table: Optional[str] = None, - other_column: Optional[str] = None, + other_column: Optional[Union[str, List[str]]] = None, ignore: bool = False, ): """ Alter the schema to mark the specified column as a foreign key to another table. - :param column: The column to mark as a foreign key. + :param column: The column to mark as a foreign key - use a list of columns + for a compound foreign key. :param other_table: The table it refers to - if omitted, will be guessed based on the column name. :param other_column: The column on the other table it - if omitted, will be guessed. + Use a list of columns for a compound foreign key. :param ignore: Set this to ``True`` to ignore an existing foreign key - otherwise a ``AlterError`` will be raised. """ - # Ensure column exists - if column not in self.columns_dict: - raise AlterError("No such column: {}".format(column)) + columns = [column] if isinstance(column, str) else list(column) + # Ensure columns exist + for col in columns: + if col not in self.columns_dict: + raise AlterError("No such column: {}".format(col)) # If other_table is not specified, attempt to guess it from the column if other_table is None: - other_table = self.guess_foreign_table(column) + if len(columns) > 1: + raise ValueError( + "other_table must be specified for a compound foreign key" + ) + other_table = self.guess_foreign_table(columns[0]) # If other_column is not specified, detect the primary key on other_table if other_column is None: - other_column = self.guess_foreign_column(other_table) + if len(columns) > 1: + other_columns = self.db.table(other_table).pks + else: + other_columns = [self.guess_foreign_column(other_table)] + elif isinstance(other_column, str): + other_columns = [other_column] + else: + other_columns = list(other_column) + if len(columns) != len(other_columns): + raise ValueError( + "Compound foreign key must have the same number of columns " + "on both sides" + ) - # Soundness check that the other column exists - if ( - not [c for c in self.db[other_table].columns if c.name == other_column] - and other_column != "rowid" - ): - raise AlterError("No such column: {}.{}".format(other_table, other_column)) + # Soundness check that the other columns exist + for other_col in other_columns: + if ( + not [c for c in self.db[other_table].columns if c.name == other_col] + and other_col != "rowid" + ): + raise AlterError("No such column: {}.{}".format(other_table, other_col)) # Check we do not already have an existing foreign key if any( fk for fk in self.foreign_keys - if fk.column == column + if fk.columns == columns and fk.other_table == other_table - and fk.other_column == other_column + and fk.other_columns == other_columns ): if ignore: return self else: raise AlterError( "Foreign key already exists for {} => {}.{}".format( - column, other_table, other_column + ", ".join(columns), other_table, ", ".join(other_columns) ) ) - self.db.add_foreign_keys([(self.name, column, other_table, other_column)]) + if len(columns) == 1: + self.db.add_foreign_keys( + [(self.name, columns[0], other_table, other_columns[0])] + ) + else: + self.db.add_foreign_keys([(self.name, columns, other_table, other_columns)]) return self def enable_counts(self) -> None: diff --git a/tests/test_foreign_keys.py b/tests/test_foreign_keys.py index d6a7071..1fabce7 100644 --- a/tests/test_foreign_keys.py +++ b/tests/test_foreign_keys.py @@ -250,3 +250,77 @@ def test_transform_drop_compound_foreign_key(compound_db, drop_foreign_keys): assert {"campus_name", "dept_code"} <= set( compound_db["courses"].columns_dict.keys() ) + + +@pytest.fixture +def courses_db(departments_db): + departments_db.create_table( + "courses", + {"course_code": str, "campus_name": str, "dept_code": str}, + pk="course_code", + ) + return departments_db + + +def test_add_compound_foreign_key(courses_db): + t = courses_db["courses"].add_foreign_key( + ["campus_name", "dept_code"], "departments", ["campus_name", "dept_code"] + ) + # Returns self + assert t.name == "courses" + fks = courses_db["courses"].foreign_keys + assert len(fks) == 1 + fk = fks[0] + assert fk.is_compound is True + assert fk.columns == ["campus_name", "dept_code"] + assert fk.other_table == "departments" + assert fk.other_columns == ["campus_name", "dept_code"] + + +def test_add_compound_foreign_key_guesses_other_columns(courses_db): + courses_db["courses"].add_foreign_key(["campus_name", "dept_code"], "departments") + fk = courses_db["courses"].foreign_keys[0] + assert fk.other_columns == ["campus_name", "dept_code"] + + +def test_add_compound_foreign_key_error_if_already_exists(courses_db): + courses_db["courses"].add_foreign_key(["campus_name", "dept_code"], "departments") + with pytest.raises(AlterError) as ex: + courses_db["courses"].add_foreign_key( + ["campus_name", "dept_code"], "departments" + ) + assert "already exists" in ex.value.args[0] + # ignore=True should not raise + courses_db["courses"].add_foreign_key( + ["campus_name", "dept_code"], "departments", ignore=True + ) + + +def test_add_compound_foreign_key_error_if_column_missing(courses_db): + with pytest.raises(AlterError): + courses_db["courses"].add_foreign_key(["campus_name", "nope"], "departments") + + +def test_db_add_foreign_keys_compound(courses_db): + courses_db.add_foreign_keys( + [ + ( + "courses", + ["campus_name", "dept_code"], + "departments", + ["campus_name", "dept_code"], + ) + ] + ) + fk = courses_db["courses"].foreign_keys[0] + assert fk.is_compound is True + assert fk.columns == ["campus_name", "dept_code"] + + +def test_index_foreign_keys_compound_creates_composite_index(compound_db): + compound_db.index_foreign_keys() + index_columns = [i.columns for i in compound_db["courses"].indexes] + assert ["campus_name", "dept_code"] in index_columns + # No separate single-column indexes for the members + assert ["campus_name"] not in index_columns + assert ["dept_code"] not in index_columns