Coverage for snekql/schema.py: 95%
114 statements
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-07 21:13 +0300
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-07 21:13 +0300
1"""SQLite schema startup for snekql table models."""
3from __future__ import annotations
5import contextlib
6from collections.abc import Sequence
7from typing import Any
9from aiosqlite import Connection, Error
11from snekql._schema_plan import (
12 PlannedColumn,
13 PlannedModel,
14 build_schema_plan,
15)
16from snekql._schema_plan import (
17 validate_schema_policy as validate_planned_schema_policy,
18)
19from snekql.errors import SchemaError, SchemaVerificationError
20from snekql.indexes import NormalizedIndex
21from snekql.model import Table
22from snekql.sqlite.identifiers import quote_identifier
23from snekql.storage import Attr, CurrentTimestamp, SchemaPolicy
24from snekql.structured_logging import ResolvedStructuredLogger
27def quote_sqlite_identifier(identifier: str) -> str:
28 """Quote a SQLite identifier with double-quote escaping."""
30 return quote_identifier(identifier)
33def _compile_column_definition(
34 name: str,
35 column: Attr[Any, Any, Any, Any, Any],
36) -> str:
37 parts = [quote_sqlite_identifier(name), column.sqlite_storage_class]
38 if column.primary_key:
39 parts.append("PRIMARY KEY")
40 if column.auto_increment:
41 parts.append("AUTOINCREMENT")
42 if column.nullable is False and not column.primary_key:
43 parts.append("NOT NULL")
44 if isinstance(column.server_default, CurrentTimestamp):
45 parts.append("DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))")
46 return " ".join(parts)
49def _compile_planned_column_definition(planned_column: PlannedColumn) -> str:
50 return _compile_column_definition(planned_column.name, planned_column.column)
53def _compile_create_table_sql(planned_model: PlannedModel) -> str:
54 column_sql = ", ".join(
55 _compile_planned_column_definition(planned_column)
56 for planned_column in planned_model.columns
57 )
58 return (
59 f"CREATE TABLE {quote_sqlite_identifier(planned_model.table_name)} "
60 f"({column_sql}) STRICT"
61 )
64def _compile_create_index_sql(table_name: str, index: NormalizedIndex) -> str:
65 unique_sql = "UNIQUE " if index.unique else ""
66 column_sql = ", ".join(
67 quote_sqlite_identifier(column_name) for column_name in index.column_names
68 )
69 return (
70 f"CREATE {unique_sql}INDEX {quote_sqlite_identifier(index.name)} "
71 f"ON {quote_sqlite_identifier(table_name)} ({column_sql})"
72 )
75def _compile_model_index_sql(planned_model: PlannedModel) -> list[str]:
76 return [
77 _compile_create_index_sql(planned_model.table_name, index)
78 for index in planned_model.indexes
79 ]
82async def _fetch_existing_create_index_sql(
83 connection: Connection,
84 table_name: str,
85) -> list[str | None]:
86 cursor = await connection.execute(
87 """
88 SELECT sql FROM sqlite_master
89 WHERE type = 'index' AND tbl_name = ?
90 ORDER BY rowid
91 """,
92 (table_name,),
93 )
94 try:
95 rows = await cursor.fetchall()
96 finally:
97 await cursor.close()
98 return [row[0] if isinstance(row[0], str) else None for row in rows]
101async def _fetch_existing_create_table_sql(
102 connection: Connection,
103 table_name: str,
104) -> str | None:
105 cursor = await connection.execute(
106 "SELECT sql FROM sqlite_master WHERE type = 'table' AND name = ?",
107 (table_name,),
108 )
109 try:
110 row = await cursor.fetchone()
111 finally:
112 await cursor.close()
113 if row is None:
114 return None
115 value = row[0]
116 if not isinstance(value, str):
117 msg = f"SQLite metadata for table {table_name!r} did not contain SQL text"
118 raise SchemaVerificationError(
119 msg,
120 )
121 return value
124def _normalize_snekql_create_table_sql(sql: str) -> str:
125 lines = [line.strip() for line in sql.strip().rstrip(";").splitlines()]
126 normalized_sql = " ".join(line for line in lines if line)
127 return normalized_sql.replace("( ", "(").replace(" )", ")").replace(" ,", ",")
130async def _report_schema_drift(
131 schema_policy: SchemaPolicy,
132 table_name: str,
133 logger: ResolvedStructuredLogger,
134) -> None:
135 message = f"schema drift detected for table {table_name!r}"
136 if schema_policy == "strict":
137 raise SchemaVerificationError(message)
138 logger.warning(
139 "schema drift detected",
140 table_name=table_name,
141 )
144async def _verify_model_indexes(
145 connection: Connection,
146 planned_model: PlannedModel,
147 schema_policy: SchemaPolicy,
148 logger: ResolvedStructuredLogger,
149) -> None:
150 expected_sql = _compile_model_index_sql(planned_model)
151 existing_sql = await _fetch_existing_create_index_sql(
152 connection,
153 planned_model.table_name,
154 )
155 if existing_sql == expected_sql:
156 logger.debug("schema indexes verified", table_name=planned_model.table_name)
157 return
158 await _report_schema_drift(schema_policy, planned_model.table_name, logger=logger)
161async def _verify_or_create_model_table(
162 connection: Connection,
163 planned_model: PlannedModel,
164 schema_policy: SchemaPolicy,
165 logger: ResolvedStructuredLogger,
166) -> None:
167 expected_sql = _compile_create_table_sql(planned_model)
168 existing_sql = await _fetch_existing_create_table_sql(
169 connection,
170 planned_model.table_name,
171 )
172 if existing_sql is None:
173 _ = await connection.execute(expected_sql)
174 logger.debug("schema table created", table_name=planned_model.table_name)
175 for index_sql in _compile_model_index_sql(planned_model):
176 _ = await connection.execute(index_sql)
177 logger.debug(
178 "schema index created",
179 table_name=planned_model.table_name,
180 sql=index_sql,
181 )
182 return
183 if _normalize_snekql_create_table_sql(
184 existing_sql,
185 ) != _normalize_snekql_create_table_sql(expected_sql):
186 await _report_schema_drift(
187 schema_policy, planned_model.table_name, logger=logger
188 )
189 return
190 logger.debug("schema table verified", table_name=planned_model.table_name)
191 await _verify_model_indexes(connection, planned_model, schema_policy, logger)
194async def _rollback_schema_setup(connection: Connection) -> None:
195 with contextlib.suppress(Error):
196 _ = await connection.execute("ROLLBACK")
199async def _initialize_sqlite_schema(
200 connection: Connection,
201 models: Sequence[type[Table[Any]]],
202 schema_policy: SchemaPolicy,
203 logger: ResolvedStructuredLogger,
204) -> None:
205 validate_planned_schema_policy(schema_policy)
206 plan = build_schema_plan(models)
207 if not plan.models:
208 return
209 logger.debug("schema startup started", model_count=len(plan.models))
210 try:
211 _ = await connection.execute("BEGIN")
212 for planned_model in plan.models:
213 await _verify_or_create_model_table(
214 connection,
215 planned_model,
216 schema_policy,
217 logger,
218 )
219 _ = await connection.execute("COMMIT")
220 logger.debug("schema startup completed", model_count=len(plan.models))
221 except Error as error:
222 await _rollback_schema_setup(connection)
223 msg = "SQLite schema setup failed"
224 raise SchemaError(msg) from error
225 except Exception:
226 await _rollback_schema_setup(connection)
227 raise
230def validate_schema_models(models: Sequence[type[Table[Any]]]) -> None:
231 """Reject duplicate resolved table names before schema startup."""
233 _ = build_schema_plan(models)
236def validate_schema_policy(schema_policy: SchemaPolicy) -> None:
237 """Reject unsupported schema policy values."""
239 validate_planned_schema_policy(schema_policy)
242async def initialize_sqlite_schema(
243 connection: Connection,
244 models: Sequence[type[Table[Any]]],
245 schema_policy: SchemaPolicy,
246 logger: ResolvedStructuredLogger,
247) -> None:
248 """Create or verify all configured SQLite tables transactionally."""
250 await _initialize_sqlite_schema(
251 connection,
252 models,
253 schema_policy,
254 logger,
255 )