Coverage for snekql/mariadb/schema.py: 92%
138 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"""MariaDB schema startup for snekql table models."""
3from __future__ import annotations
5from collections.abc import Sequence
6from dataclasses import dataclass
7from typing import Any, cast
9from snekql._schema_plan import (
10 PlannedColumn,
11 PlannedModel,
12 build_schema_plan,
13)
14from snekql._schema_plan import (
15 validate_schema_policy as validate_planned_schema_policy,
16)
17from snekql.errors import SchemaError, SchemaVerificationError
18from snekql.indexes import NormalizedIndex
19from snekql.mariadb.identifiers import quote_identifier
20from snekql.model import Table
21from snekql.storage import Attr, CurrentTimestamp, SchemaPolicy
22from snekql.structured_logging import ResolvedStructuredLogger
25@dataclass(frozen=True)
26class _ColumnSignature:
27 """Normalized MariaDB column metadata used for drift verification."""
29 auto_increment: bool
30 data_type: str
31 max_length: int | None
32 name: str
33 nullable: bool
34 primary_key: bool
37@dataclass(frozen=True)
38class _IndexSignature:
39 """Normalized MariaDB index metadata used for drift verification."""
41 column_names: tuple[str, ...]
42 name: str
43 unique: bool
46def _compile_column_type(column: Attr[Any, Any, Any, Any, Any]) -> str:
47 """Map the initial shared value families to MariaDB column types."""
49 column_types = {
50 "Blob": "BLOB",
51 "Boolean": "BOOLEAN",
52 "DateTime": "DATETIME(3)",
53 "Integer": "BIGINT",
54 "Json": "JSON",
55 "Real": "DOUBLE",
56 "Text": "VARCHAR(255)",
57 }
58 try:
59 return column_types[column.storage_type_name]
60 except KeyError as error:
61 msg = f"unsupported MariaDB column type: {column.storage_type_name}"
62 raise SchemaError(msg) from error
65def _column_data_type(column: Attr[Any, Any, Any, Any, Any]) -> str:
66 """Return information_schema.DATA_TYPE expected for a column."""
68 data_types = {
69 "Blob": "blob",
70 "Boolean": "tinyint",
71 "DateTime": "datetime",
72 "Integer": "bigint",
73 "Json": "longtext",
74 "Real": "double",
75 "Text": "varchar",
76 }
77 try:
78 return data_types[column.storage_type_name]
79 except KeyError as error:
80 msg = f"unsupported MariaDB column type: {column.storage_type_name}"
81 raise SchemaError(msg) from error
84def _column_max_length(column: Attr[Any, Any, Any, Any, Any]) -> int | None:
85 if column.storage_type_name == "Text":
86 return 255
87 return None
90def _expected_column_signature(planned_column: PlannedColumn) -> _ColumnSignature:
91 column = planned_column.column
92 return _ColumnSignature(
93 auto_increment=column.auto_increment,
94 data_type=_column_data_type(column),
95 max_length=_column_max_length(column),
96 name=planned_column.name,
97 nullable=column.nullable is not False and not column.primary_key,
98 primary_key=column.primary_key,
99 )
102def _compile_column_definition(
103 name: str,
104 column: Attr[Any, Any, Any, Any, Any],
105) -> str:
106 parts = [quote_identifier(name), _compile_column_type(column)]
107 if column.nullable is False or column.primary_key:
108 parts.append("NOT NULL")
109 if column.auto_increment:
110 parts.append("AUTO_INCREMENT")
111 if column.primary_key:
112 parts.append("PRIMARY KEY")
113 if isinstance(column.server_default, CurrentTimestamp):
114 parts.append("DEFAULT CURRENT_TIMESTAMP(3)")
115 return " ".join(parts)
118def _compile_planned_column_definition(planned_column: PlannedColumn) -> str:
119 return _compile_column_definition(planned_column.name, planned_column.column)
122def _compile_create_table_sql(planned_model: PlannedModel) -> str:
123 column_sql = ", ".join(
124 _compile_planned_column_definition(planned_column)
125 for planned_column in planned_model.columns
126 )
127 return f"CREATE TABLE {quote_identifier(planned_model.table_name)} ({column_sql})"
130def _compile_create_index_sql(table_name: str, index: NormalizedIndex) -> str:
131 unique_sql = "UNIQUE " if index.unique else ""
132 column_sql = ", ".join(
133 quote_identifier(column_name) for column_name in index.column_names
134 )
135 return (
136 f"CREATE {unique_sql}INDEX {quote_identifier(index.name)} "
137 f"ON {quote_identifier(table_name)} ({column_sql})"
138 )
141def _expected_index_signatures(planned_model: PlannedModel) -> list[_IndexSignature]:
142 return [
143 _IndexSignature(
144 column_names=index.column_names,
145 name=index.name,
146 unique=index.unique,
147 )
148 for index in planned_model.indexes
149 ]
152async def _close_cursor(cursor: object) -> None:
153 close_result = cast("Any", cursor).close()
154 if close_result is not None:
155 _ = await close_result
158async def _execute(
159 connection: object,
160 sql: str,
161 params: tuple[object, ...] = (),
162) -> None:
163 """Execute one MariaDB schema statement with a dynamically imported driver."""
165 cursor = await cast("Any", connection).cursor()
166 try:
167 _ = await cursor.execute(sql, params)
168 finally:
169 await _close_cursor(cursor)
172async def _fetchall(
173 connection: object,
174 sql: str,
175 params: tuple[object, ...] = (),
176) -> Sequence[Sequence[object]]:
177 cursor = await cast("Any", connection).cursor()
178 try:
179 _ = await cursor.execute(sql, params)
180 rows = await cursor.fetchall()
181 finally:
182 await _close_cursor(cursor)
183 return [cast("Sequence[object]", row) for row in rows]
186async def _table_exists(connection: object, table_name: str) -> bool:
187 rows = await _fetchall(
188 connection,
189 """
190 SELECT 1
191 FROM INFORMATION_SCHEMA.TABLES
192 WHERE TABLE_SCHEMA = DATABASE()
193 AND TABLE_NAME = %s
194 """,
195 (table_name,),
196 )
197 return bool(rows)
200async def _fetch_existing_column_signatures(
201 connection: object,
202 table_name: str,
203) -> list[_ColumnSignature]:
204 rows = await _fetchall(
205 connection,
206 """
207 SELECT COLUMN_NAME, DATA_TYPE, CHARACTER_MAXIMUM_LENGTH, IS_NULLABLE,
208 COLUMN_KEY, EXTRA
209 FROM INFORMATION_SCHEMA.COLUMNS
210 WHERE TABLE_SCHEMA = DATABASE()
211 AND TABLE_NAME = %s
212 ORDER BY ORDINAL_POSITION
213 """,
214 (table_name,),
215 )
216 signatures: list[_ColumnSignature] = []
217 for row in rows:
218 name, data_type, max_length, nullable, column_key, extra = row
219 parsed_max_length = (
220 int(max_length) if isinstance(max_length, int | str) else None
221 )
222 signatures.append(
223 _ColumnSignature(
224 auto_increment="auto_increment" in str(extra),
225 data_type=str(data_type),
226 max_length=parsed_max_length,
227 name=str(name),
228 nullable=nullable == "YES",
229 primary_key=column_key == "PRI",
230 )
231 )
232 return signatures
235async def _fetch_existing_index_signatures(
236 connection: object,
237 table_name: str,
238) -> list[_IndexSignature]:
239 rows = await _fetchall(
240 connection,
241 """
242 SELECT INDEX_NAME, NON_UNIQUE,
243 GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX)
244 FROM INFORMATION_SCHEMA.STATISTICS
245 WHERE TABLE_SCHEMA = DATABASE()
246 AND TABLE_NAME = %s
247 AND INDEX_NAME <> 'PRIMARY'
248 GROUP BY INDEX_NAME, NON_UNIQUE
249 ORDER BY INDEX_NAME
250 """,
251 (table_name,),
252 )
253 indexes: list[_IndexSignature] = []
254 for row in rows:
255 name, non_unique, column_csv = row
256 indexes.append(
257 _IndexSignature(
258 column_names=tuple(str(column_csv).split(",")),
259 name=str(name),
260 unique=non_unique == 0,
261 )
262 )
263 return indexes
266async def _report_schema_drift(
267 schema_policy: SchemaPolicy,
268 table_name: str,
269 logger: ResolvedStructuredLogger,
270) -> None:
271 message = f"schema drift detected for table {table_name!r}"
272 if schema_policy == "strict":
273 raise SchemaVerificationError(message)
274 logger.warning(
275 "schema drift detected",
276 table_name=table_name,
277 )
280async def _verify_model_schema(
281 connection: object,
282 planned_model: PlannedModel,
283 schema_policy: SchemaPolicy,
284 logger: ResolvedStructuredLogger,
285) -> None:
286 expected_columns = [
287 _expected_column_signature(planned_column)
288 for planned_column in planned_model.columns
289 ]
290 existing_columns = await _fetch_existing_column_signatures(
291 connection,
292 planned_model.table_name,
293 )
294 if existing_columns != expected_columns:
295 await _report_schema_drift(
296 schema_policy, planned_model.table_name, logger=logger
297 )
298 return
299 expected_indexes = sorted(
300 _expected_index_signatures(planned_model), key=lambda index: index.name
301 )
302 existing_indexes = await _fetch_existing_index_signatures(
303 connection,
304 planned_model.table_name,
305 )
306 if existing_indexes != expected_indexes:
307 await _report_schema_drift(
308 schema_policy, planned_model.table_name, logger=logger
309 )
310 return
311 logger.debug("schema table verified", table_name=planned_model.table_name)
312 logger.debug("schema indexes verified", table_name=planned_model.table_name)
315async def _create_model_schema(
316 connection: object,
317 planned_model: PlannedModel,
318 logger: ResolvedStructuredLogger,
319) -> None:
320 await _execute(connection, _compile_create_table_sql(planned_model))
321 logger.debug("schema table created", table_name=planned_model.table_name)
322 for index in planned_model.indexes:
323 sql = _compile_create_index_sql(planned_model.table_name, index)
324 await _execute(connection, sql)
325 logger.debug(
326 "schema index created",
327 table_name=planned_model.table_name,
328 sql=sql,
329 )
332async def initialize_mariadb_schema(
333 connection: object,
334 models: Sequence[type[Table[Any]]],
335 schema_policy: SchemaPolicy,
336 logger: ResolvedStructuredLogger,
337) -> None:
338 """Create or verify all configured MariaDB tables."""
340 validate_planned_schema_policy(schema_policy)
341 plan = build_schema_plan(models)
342 if not plan.models:
343 return
344 logger.debug("schema startup started", model_count=len(plan.models))
345 for planned_model in plan.models:
346 if await _table_exists(connection, planned_model.table_name):
347 await _verify_model_schema(
348 connection,
349 planned_model,
350 schema_policy,
351 logger,
352 )
353 else:
354 await _create_model_schema(connection, planned_model, logger)
355 logger.debug("schema startup completed", model_count=len(plan.models))