table.checks, table.column_checks, table.table_checks, closes #834

Refs #762
This commit is contained in:
Simon Willison 2026-08-11 21:58:00 -07:00
commit 3db0c57a3b
6 changed files with 775 additions and 1 deletions

View file

@ -9,6 +9,7 @@
Unreleased
----------
- New ``table.checks``, ``table.column_checks`` and ``table.table_checks`` introspection properties expose column-level and table-level ``CHECK`` constraints. (:issue:`834`)
- ``table.transform()`` now works for tables that are referenced by views. Previously the ``ALTER TABLE ... RENAME TO`` step raised ``no such table`` if a view referenced the table being transformed. View definitions are left unchanged - see :ref:`python_api_transform_views`. This also fixes a bug where ``transform(keep_table=...)`` silently rewrote dependent views to point at the frozen backup table instead of the live one. (:issue:`831`)
.. _v3_39_1:

View file

@ -2480,6 +2480,43 @@ Almost all SQLite tables have a ``rowid`` column, but a table with no explicitly
False
.. _python_api_introspection_checks:
.checks
-------
The ``.checks`` property returns the column-level and table-level ``CHECK`` constraints defined on a table, as a list of ``Check`` objects. Each object has ``check`` (the expression inside ``CHECK (...)``), ``name``, ``column`` and ``options`` attributes. ``column`` is an empty string for a table-level check. ``options`` contains a list of values only when a column check consists entirely of ``column IN (literal, ...)``. The original constraint fragment is available as ``sql``; ``start`` and ``end`` are its offsets within ``table.schema``.
.. code-block:: python
>>> db["scores"].checks
[Check(check='score > 0', name='positive', column='score', options=None),
Check(check='score <= maximum', name='within_maximum', column='', options=None)]
.. _python_api_introspection_column_checks:
.column_checks
--------------
The ``.column_checks`` property returns the column-level checks grouped by column name:
.. code-block:: python
>>> db["scores"].column_checks
{'score': [Check(check='score > 0', name='positive', column='score', options=None)]}
.. _python_api_introspection_table_checks:
.table_checks
-------------
The ``.table_checks`` property returns only the table-level checks:
.. code-block:: python
>>> db["scores"].table_checks
[Check(check='score <= maximum', name='within_maximum', column='', options=None)]
.. _python_api_introspection_foreign_keys:
.foreign_keys

View file

@ -0,0 +1,551 @@
"""Helpers for parsing CHECK constraints from SQLite CREATE TABLE SQL.
SQLite does not expose CHECK constraints through a pragma, so preserving them
across a table rebuild requires reading ``sqlite_schema.sql``. This module is
deliberately small, but it uses a real lexer: strings, quoted identifiers and
comments are opaque, every token retains its source span and malformed input is
reported instead of being silently under-parsed.
"""
import re
from dataclasses import dataclass, field
from typing import Any
@dataclass
class Check:
check: str
name: str = ""
column: str = ""
options: list[Any] | None = None
# Source details are excluded from equality and repr so callers can compare
# semantic constraints while still having the original SQL available for
# diagnostics or future lossless edits.
sql: str = field(default="", compare=False, repr=False)
start: int = field(default=-1, compare=False, repr=False)
end: int = field(default=-1, compare=False, repr=False)
class ParseError(ValueError):
pass
@dataclass(frozen=True)
class _Token:
kind: str
text: str
start: int
end: int
def is_keyword(self, keyword: str) -> bool:
return self.kind == "word" and self.text.upper() == keyword
_PUNCTUATION = frozenset("(),.;+-*/%<>=!~|&?:")
_TRIVIA = frozenset(("whitespace", "comment"))
_TABLE_CONSTRAINT_KEYWORDS = frozenset(("PRIMARY", "UNIQUE", "CHECK", "FOREIGN"))
_OTHER_COLUMN_CONSTRAINT_KEYWORDS = frozenset(
("PRIMARY", "UNIQUE", "REFERENCES", "DEFAULT", "NOT", "COLLATE", "GENERATED")
)
_SQLITE_KEYWORDS = frozenset(
(
"ABORT",
"ACTION",
"ADD",
"AFTER",
"ALL",
"ALTER",
"ANALYZE",
"AND",
"AS",
"ASC",
"ATTACH",
"AUTOINCREMENT",
"BEFORE",
"BEGIN",
"BETWEEN",
"BY",
"CASCADE",
"CASE",
"CAST",
"CHECK",
"COLLATE",
"COLUMN",
"COMMIT",
"CONFLICT",
"CONSTRAINT",
"CREATE",
"CROSS",
"CURRENT_DATE",
"CURRENT_TIME",
"CURRENT_TIMESTAMP",
"DATABASE",
"DEFAULT",
"DEFERRABLE",
"DEFERRED",
"DELETE",
"DESC",
"DETACH",
"DISTINCT",
"DO",
"DROP",
"EACH",
"ELSE",
"END",
"ESCAPE",
"EXCEPT",
"EXCLUDE",
"EXCLUSIVE",
"EXISTS",
"EXPLAIN",
"FAIL",
"FALSE",
"FILTER",
"FIRST",
"FOLLOWING",
"FOR",
"FOREIGN",
"FROM",
"FULL",
"GENERATED",
"GLOB",
"GROUP",
"GROUPS",
"HAVING",
"IF",
"IGNORE",
"IMMEDIATE",
"IN",
"INDEX",
"INDEXED",
"INITIALLY",
"INNER",
"INSERT",
"INSTEAD",
"INTERSECT",
"INTO",
"IS",
"ISNULL",
"JOIN",
"KEY",
"LAST",
"LEFT",
"LIKE",
"LIMIT",
"MATCH",
"MATERIALIZED",
"NATURAL",
"NO",
"NOT",
"NOTHING",
"NOTNULL",
"NULL",
"NULLS",
"OF",
"OFFSET",
"ON",
"OR",
"ORDER",
"OTHERS",
"OUTER",
"OVER",
"PARTITION",
"PLAN",
"PRAGMA",
"PRECEDING",
"PRIMARY",
"QUERY",
"RAISE",
"RANGE",
"RECURSIVE",
"REFERENCES",
"REGEXP",
"REINDEX",
"RELEASE",
"RENAME",
"REPLACE",
"RESTRICT",
"RETURNING",
"RIGHT",
"ROLLBACK",
"ROW",
"ROWS",
"SAVEPOINT",
"SELECT",
"SET",
"STRICT",
"TABLE",
"TEMP",
"TEMPORARY",
"THEN",
"TIES",
"TO",
"TRANSACTION",
"TRIGGER",
"TRUE",
"UNBOUNDED",
"UNION",
"UNIQUE",
"UPDATE",
"USING",
"VACUUM",
"VALUES",
"VIEW",
"VIRTUAL",
"WHEN",
"WHERE",
"WINDOW",
"WITH",
"WITHOUT",
)
)
_INTEGER_RE = re.compile(r"[+-]?(?:0[xX][0-9a-fA-F]+|[0-9]+)\Z")
_FLOAT_RE = re.compile(
r"[+-]?(?:(?:[0-9]+\.[0-9]*|\.[0-9]+)(?:[eE][+-]?[0-9]+)?|"
r"[0-9]+[eE][+-]?[0-9]+)\Z"
)
def _lex(sql: str) -> list[_Token]:
tokens: list[_Token] = []
i = 0
while i < len(sql):
start = i
char = sql[i]
if char.isspace():
i += 1
while i < len(sql) and sql[i].isspace():
i += 1
tokens.append(_Token("whitespace", sql[start:i], start, i))
continue
if sql.startswith("--", i):
newline = sql.find("\n", i + 2)
i = len(sql) if newline == -1 else newline + 1
tokens.append(_Token("comment", sql[start:i], start, i))
continue
if sql.startswith("/*", i):
end = sql.find("*/", i + 2)
if end == -1:
raise ParseError("Unterminated SQL comment")
i = end + 2
tokens.append(_Token("comment", sql[start:i], start, i))
continue
if char in ("'", '"', "`"):
quote = char
i += 1
while i < len(sql):
if sql[i] == quote:
if i + 1 < len(sql) and sql[i + 1] == quote:
i += 2
continue
i += 1
break
i += 1
else:
raise ParseError(f"Unterminated {quote} quoted token")
kind = "string" if quote == "'" else "identifier"
tokens.append(_Token(kind, sql[start:i], start, i))
continue
if char == "[":
end = sql.find("]", i + 1)
if end == -1:
raise ParseError("Unterminated [ quoted identifier")
i = end + 1
tokens.append(_Token("identifier", sql[start:i], start, i))
continue
if char in _PUNCTUATION:
i += 1
tokens.append(_Token("punct", char, start, i))
continue
# SQLite accepts any character >= U+0080 in a bare identifier. More
# generally, consume until a lexical delimiter rather than relying on
# Python's narrower definition of an alphanumeric character.
i += 1
while i < len(sql):
if sql[i].isspace() or sql[i] in _PUNCTUATION or sql[i] in "'\"`[":
break
i += 1
tokens.append(_Token("word", sql[start:i], start, i))
return tokens
def _meaningful(tokens: list[_Token]) -> list[_Token]:
return [token for token in tokens if token.kind not in _TRIVIA]
def _unquote(token: str) -> str:
if len(token) >= 2 and token[0] in ("'", '"', "`") and token[-1] == token[0]:
return token[1:-1].replace(token[0] * 2, token[0])
if len(token) >= 2 and token[0] == "[" and token[-1] == "]":
return token[1:-1]
return token
def _matching_paren(tokens: list[_Token], open_index: int) -> int:
if tokens[open_index].text != "(":
raise ParseError("Expected an opening parenthesis")
depth = 0
for index in range(open_index, len(tokens)):
if tokens[index].text == "(":
depth += 1
elif tokens[index].text == ")":
depth -= 1
if depth == 0:
return index
raise ParseError("Unbalanced parentheses")
def _split_spans(sql: str, tokens: list[_Token]) -> list[tuple[str, int, int]]:
if not tokens:
return []
items: list[tuple[str, int, int]] = []
depth = 0
start = tokens[0].start
for token in tokens:
if token.text == "(":
depth += 1
elif token.text == ")":
depth -= 1
if depth < 0:
raise ParseError("Unbalanced parentheses")
elif token.text == "," and depth == 0:
raw = sql[start : token.start]
item = raw.strip()
if item:
item_start = start + len(raw) - len(raw.lstrip())
items.append((item, item_start, item_start + len(item)))
start = token.end
if depth:
raise ParseError("Unbalanced parentheses")
raw = sql[start : tokens[-1].end]
item = raw.strip()
if item:
item_start = start + len(raw) - len(raw.lstrip())
items.append((item, item_start, item_start + len(item)))
return items
def _split_ranges(sql: str, tokens: list[_Token]) -> list[str]:
return [item for item, _, _ in _split_spans(sql, tokens)]
def _strip_outer_parens(tokens: list[_Token]) -> list[_Token]:
while tokens and tokens[0].text == "(":
close = _matching_paren(tokens, 0)
if close != len(tokens) - 1:
break
tokens = tokens[1:-1]
return tokens
_NO_LITERAL = object()
def _literal_value(text: str) -> Any:
tokens = _meaningful(_lex(text))
if len(tokens) == 1 and tokens[0].kind == "string":
return _unquote(tokens[0].text)
raw = "".join(token.text for token in tokens)
if raw.upper() == "NULL":
return None
if raw.upper() == "TRUE":
return True
if raw.upper() == "FALSE":
return False
if _INTEGER_RE.fullmatch(raw):
try:
return (
int(raw, 16) if raw.lower().lstrip("+-").startswith("0x") else int(raw)
)
except ValueError:
return _NO_LITERAL
if _FLOAT_RE.fullmatch(raw):
try:
return float(raw)
except ValueError:
return _NO_LITERAL
return _NO_LITERAL
def _ascii_fold(identifier: str) -> str:
return identifier.translate(
str.maketrans("ABCDEFGHIJKLMNOPQRSTUVWXYZ", "abcdefghijklmnopqrstuvwxyz")
)
def _parse_options(expression: str, column: str) -> list[Any] | None:
tokens = _strip_outer_parens(_meaningful(_lex(expression)))
if len(tokens) < 4:
return None
lhs = tokens[0]
if lhs.kind not in ("word", "identifier"):
return None
if column and _ascii_fold(_unquote(lhs.text)) != _ascii_fold(column):
return None
if not tokens[1].is_keyword("IN") or tokens[2].text != "(":
return None
close = _matching_paren(tokens, 2)
if close != len(tokens) - 1:
return None
inner = expression[tokens[2].end : tokens[close].start]
inner_tokens = _lex(inner)
if not _meaningful(inner_tokens):
return []
values = []
for item in _split_ranges(inner, inner_tokens):
value = _literal_value(item)
if value is _NO_LITERAL:
return None
values.append(value)
return values
def _check_after(
item: str,
tokens: list[_Token],
check_index: int,
name: str,
column: str,
constraint_start: int,
base_offset: int,
) -> tuple[Check, int]:
if check_index + 1 >= len(tokens) or tokens[check_index + 1].text != "(":
raise ParseError("CHECK must be followed by a parenthesized expression")
close = _matching_paren(tokens, check_index + 1)
expression = item[tokens[check_index + 1].end : tokens[close].start].strip()
source_start = tokens[constraint_start].start
source_end = tokens[close].end
return (
Check(
expression,
name=name,
column=column,
options=_parse_options(expression, column),
sql=item[source_start:source_end],
start=base_offset + source_start,
end=base_offset + source_end,
),
close + 1,
)
def _column_checks(
item: str, tokens: list[_Token], column: str, base_offset: int
) -> list[Check]:
checks: list[Check] = []
pending_name = ""
pending_start: int | None = None
index = 1
while index < len(tokens):
token = tokens[index]
if token.text == "(":
index = _matching_paren(tokens, index) + 1
continue
if token.is_keyword("CONSTRAINT"):
if index + 1 >= len(tokens):
raise ParseError("CONSTRAINT is missing its name")
pending_name = _unquote(tokens[index + 1].text)
pending_start = index
index += 2
continue
if token.is_keyword("CHECK"):
check, index = _check_after(
item,
tokens,
index,
pending_name,
column,
pending_start if pending_start is not None else index,
base_offset,
)
checks.append(check)
pending_name = ""
pending_start = None
continue
if (
token.kind == "word"
and token.text.upper() in _OTHER_COLUMN_CONSTRAINT_KEYWORDS
):
pending_name = ""
pending_start = None
index += 1
return checks
def parse_checks(create_sql: str) -> list[Check]:
"""Return CHECK constraints from a valid SQLite CREATE TABLE statement."""
all_tokens = _lex(create_sql)
tokens = _meaningful(all_tokens)
if not tokens or not tokens[0].is_keyword("CREATE"):
raise ParseError("Expected CREATE TABLE")
index = 1
if index < len(tokens) and (
tokens[index].is_keyword("TEMP") or tokens[index].is_keyword("TEMPORARY")
):
index += 1
if index < len(tokens) and tokens[index].is_keyword("VIRTUAL"):
return []
if index >= len(tokens) or not tokens[index].is_keyword("TABLE"):
raise ParseError("Expected CREATE TABLE")
index += 1
if (
index + 2 < len(tokens)
and tokens[index].is_keyword("IF")
and tokens[index + 1].is_keyword("NOT")
and tokens[index + 2].is_keyword("EXISTS")
):
index += 3
if index >= len(tokens):
raise ParseError("CREATE TABLE is missing its table name")
index += 1
if index + 1 < len(tokens) and tokens[index].text == ".":
index += 2
if index < len(tokens) and tokens[index].is_keyword("AS"):
return []
if index >= len(tokens) or tokens[index].text != "(":
raise ParseError("CREATE TABLE is missing its column list")
close = _matching_paren(tokens, index)
trailing = tokens[close + 1 :]
allowed_trailing = {"STRICT", "WITHOUT", "ROWID", ",", ";"}
if any(token.text.upper() not in allowed_trailing for token in trailing):
raise ParseError("Unexpected SQL after CREATE TABLE column list")
body_start = tokens[index].end
body_end = tokens[close].start
body = create_sql[body_start:body_end]
body_tokens = _lex(body)
checks: list[Check] = []
for item, item_start, _ in _split_spans(body, body_tokens):
item_tokens = _meaningful(_lex(item))
if not item_tokens:
continue
item_index = 0
constraint_name = ""
if item_tokens[item_index].is_keyword("CONSTRAINT"):
if len(item_tokens) < 2:
raise ParseError("CONSTRAINT is missing its name")
constraint_name = _unquote(item_tokens[1].text)
item_index = 2
head = item_tokens[item_index] if item_index < len(item_tokens) else None
if (
head
and head.kind == "word"
and head.text.upper() in _TABLE_CONSTRAINT_KEYWORDS
):
if head.is_keyword("CHECK"):
check, _ = _check_after(
item,
item_tokens,
item_index,
constraint_name,
"",
0,
body_start + item_start,
)
checks.append(check)
continue
column = _unquote(item_tokens[0].text)
checks.extend(
_column_checks(item, item_tokens, column, body_start + item_start)
)
return checks

View file

@ -27,6 +27,7 @@ from typing_extensions import Self
from sqlite_utils.plugins import ensure_plugins_loaded, pm
from .create_table_parser import Check, parse_checks
from .utils import (
OperationalError,
chunks,
@ -2196,6 +2197,27 @@ class Table(Queryable):
"Does this table use ``rowid`` for its primary key (no other primary keys are specified)?"
return not any(column for column in self.columns if column.is_pk)
@property
def checks(self) -> list[Check]:
"List of column-level and table-level CHECK constraints on this table."
if not self.exists() or self.virtual_table_using is not None:
return []
return parse_checks(self.schema)
@property
def column_checks(self) -> dict[str, list[Check]]:
"CHECK constraints grouped by the column on which they are defined."
checks: dict[str, list[Check]] = {}
for check in self.checks:
if check.column:
checks.setdefault(check.column, []).append(check)
return checks
@property
def table_checks(self) -> list[Check]:
"Table-level CHECK constraints on this table."
return [check for check in self.checks if not check.column]
def get(self, pk_values: list | tuple | str | int) -> dict:
"""
Return row (as dictionary) for the specified primary key.

View file

@ -0,0 +1,138 @@
import sqlite3
import hypothesis.strategies as st
import pytest
from hypothesis import given
from sqlite_utils.create_table_parser import Check, ParseError, parse_checks
def test_parse_column_and_table_checks():
sql = """
CREATE TABLE people (
age INTEGER CONSTRAINT positive CHECK (age > 0),
status TEXT CHECK(status IN ('active', 'inactive')),
CONSTRAINT adult CHECK(age >= 18)
)
"""
assert parse_checks(sql) == [
Check("age > 0", name="positive", column="age"),
Check(
"status IN ('active', 'inactive')",
column="status",
options=["active", "inactive"],
),
Check("age >= 18", name="adult"),
]
checks = parse_checks(sql)
assert checks[0].sql == "CONSTRAINT positive CHECK (age > 0)"
assert sql[checks[0].start : checks[0].end] == checks[0].sql
assert checks[1].sql == "CHECK(status IN ('active', 'inactive'))"
assert sql[checks[2].start : checks[2].end] == checks[2].sql
def test_comments_are_trivia_not_constraints():
sql = """
CREATE /* fake CHECK (nope), ( */ TABLE t (
a INTEGER /* CHECK (a < 0), phantom */,
b INTEGER CHECK /* between keyword and expression */ (b > 0),
/* CHECK (also_fake) */ CONSTRAINT upper CHECK(b < 10)
)
"""
sqlite3.connect(":memory:").execute(sql)
assert parse_checks(sql) == [
Check("b > 0", column="b"),
Check("b < 10", name="upper"),
]
@pytest.mark.parametrize(
"expression,expected",
[
("value IN ('one', 'two')", ["one", "two"]),
("((value IN ('one', 'two')))", ["one", "two"]),
("value NOT IN ('one', 'two')", None),
("value IN ('one', 'two') OR enabled", None),
("other IN ('one', 'two')", None),
("value IN (lower('one'), 'two')", None),
('value IN ("other")', None),
],
)
def test_options_only_for_exact_literal_in_check(expression, expected):
sql = f"CREATE TABLE t(value TEXT CHECK({expression}), enabled INTEGER, other TEXT)"
sqlite3.connect(":memory:").execute(sql)
assert parse_checks(sql)[0].options == expected
@pytest.mark.parametrize("column", ["💩x", "e\u0301"])
def test_unquoted_unicode_identifiers(column):
sql = f"CREATE TABLE t({column} INTEGER CHECK({column} > 0))"
sqlite3.connect(":memory:").execute(sql)
assert parse_checks(sql) == [Check(f"{column} > 0", column=column)]
@pytest.mark.parametrize(
"sql",
[
"SELECT CHECK(x > 0)",
"CREATE TABLE t(x INTEGER CHECK(x > 0)",
"CREATE TABLE t(x TEXT CHECK(x != 'unterminated))",
"CREATE TABLE t(x INTEGER /* unterminated)",
],
)
def test_invalid_sql_raises_parse_error(sql):
with pytest.raises(ParseError):
parse_checks(sql)
def test_virtual_table_has_no_checks():
assert (
parse_checks("CREATE /* comment */ VIRTUAL TABLE search USING fts5(text)") == []
)
comment_or_space = st.sampled_from(
[
" ",
"\n ",
"/* comment with , ( ) and CHECK(fake) */",
"-- comment with , ( ) and CHECK(fake)\n",
]
)
@given(gaps=st.lists(comment_or_space, min_size=5, max_size=5))
def test_comments_and_whitespace_can_separate_check_tokens(gaps):
sql = (
f"CREATE{gaps[0]}TABLE{gaps[1]}t{gaps[2]}("
f"value INTEGER CHECK{gaps[3]}(value{gaps[4]}> 0))"
)
connection = sqlite3.connect(":memory:")
connection.execute(sql)
stored_sql = connection.execute(
"select sql from sqlite_schema where name = 't'"
).fetchone()[0]
assert parse_checks(stored_sql) == [Check(f"value{gaps[4]}> 0", column="value")]
safe_string_text = st.text(
alphabet=st.characters(
blacklist_categories=("Cc", "Cs"),
blacklist_characters=("'",),
),
max_size=40,
)
@given(value=safe_string_text)
def test_check_like_text_inside_strings_is_opaque(value):
sql = f"CREATE TABLE t(value TEXT CHECK(value != '{value}'))"
connection = sqlite3.connect(":memory:")
connection.execute(sql)
stored_sql = connection.execute(
"select sql from sqlite_schema where name = 't'"
).fetchone()[0]
checks = parse_checks(stored_sql)
assert len(checks) == 1
assert checks[0].column == "value"
assert checks[0].check == f"value != '{value}'"

View file

@ -1,6 +1,6 @@
import pytest
from sqlite_utils.db import Database, Index, View, XIndex, XIndexColumn
from sqlite_utils.db import Check, Database, Index, View, XIndex, XIndexColumn
def _check_supports_strict():
@ -177,6 +177,31 @@ def test_pks(fresh_db, pk, expected):
assert expected == fresh_db["foo"].pks
def test_checks(fresh_db):
fresh_db.execute("""
CREATE TABLE scores (
score INTEGER CONSTRAINT positive CHECK(score > 0),
maximum INTEGER,
CONSTRAINT within_maximum CHECK(score <= maximum)
)
""")
scores = fresh_db["scores"]
expected_column = Check("score > 0", name="positive", column="score")
expected_table = Check("score <= maximum", name="within_maximum")
assert scores.checks == [expected_column, expected_table]
assert scores.column_checks == {"score": [expected_column]}
assert scores.table_checks == [expected_table]
assert scores.checks[0].sql == "CONSTRAINT positive CHECK(score > 0)"
def test_checks_nonexistent_and_virtual_tables(fresh_db):
assert fresh_db["does_not_exist"].checks == []
fresh_db["searchable"].insert({"text": "hello"}).enable_fts(
["text"], fts_version="FTS5"
)
assert fresh_db["searchable_fts"].checks == []
def test_triggers_and_triggers_dict(fresh_db):
assert [] == fresh_db.triggers
authors = fresh_db["authors"]