From 3db0c57a3bc9d8468db430ebe0ffd0da213fdda3 Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Tue, 11 Aug 2026 21:58:00 -0700 Subject: [PATCH] table.checks, table.column_checks, table.table_checks, closes #834 Refs #762 --- docs/changelog.rst | 1 + docs/python-api.rst | 37 ++ sqlite_utils/create_table_parser.py | 551 ++++++++++++++++++++++++++++ sqlite_utils/db.py | 22 ++ tests/test_create_table_parser.py | 138 +++++++ tests/test_introspect.py | 27 +- 6 files changed, 775 insertions(+), 1 deletion(-) create mode 100644 sqlite_utils/create_table_parser.py create mode 100644 tests/test_create_table_parser.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 4ae53a6..77f644c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -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: diff --git a/docs/python-api.rst b/docs/python-api.rst index 53a47dd..93f8a3d 100644 --- a/docs/python-api.rst +++ b/docs/python-api.rst @@ -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 diff --git a/sqlite_utils/create_table_parser.py b/sqlite_utils/create_table_parser.py new file mode 100644 index 0000000..9c0a9aa --- /dev/null +++ b/sqlite_utils/create_table_parser.py @@ -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 diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 5307931..d85ca41 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -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. diff --git a/tests/test_create_table_parser.py b/tests/test_create_table_parser.py new file mode 100644 index 0000000..e7ab3e8 --- /dev/null +++ b/tests/test_create_table_parser.py @@ -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}'" diff --git a/tests/test_introspect.py b/tests/test_introspect.py index b7e8fc2..2a8d579 100644 --- a/tests/test_introspect.py +++ b/tests/test_introspect.py @@ -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"]