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 <noreply@anthropic.com>
This commit is contained in:
Simon Willison 2026-07-05 14:37:32 -07:00
commit b75edf4b30
3 changed files with 177 additions and 47 deletions

View file

@ -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

View file

@ -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:

View file

@ -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