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

1"""MariaDB schema startup for snekql table models.""" 

2 

3from __future__ import annotations 

4 

5from collections.abc import Sequence 

6from dataclasses import dataclass 

7from typing import Any, cast 

8 

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 

23 

24 

25@dataclass(frozen=True) 

26class _ColumnSignature: 

27 """Normalized MariaDB column metadata used for drift verification.""" 

28 

29 auto_increment: bool 

30 data_type: str 

31 max_length: int | None 

32 name: str 

33 nullable: bool 

34 primary_key: bool 

35 

36 

37@dataclass(frozen=True) 

38class _IndexSignature: 

39 """Normalized MariaDB index metadata used for drift verification.""" 

40 

41 column_names: tuple[str, ...] 

42 name: str 

43 unique: bool 

44 

45 

46def _compile_column_type(column: Attr[Any, Any, Any, Any, Any]) -> str: 

47 """Map the initial shared value families to MariaDB column types.""" 

48 

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 

63 

64 

65def _column_data_type(column: Attr[Any, Any, Any, Any, Any]) -> str: 

66 """Return information_schema.DATA_TYPE expected for a column.""" 

67 

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 

82 

83 

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 

88 

89 

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 ) 

100 

101 

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) 

116 

117 

118def _compile_planned_column_definition(planned_column: PlannedColumn) -> str: 

119 return _compile_column_definition(planned_column.name, planned_column.column) 

120 

121 

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})" 

128 

129 

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 ) 

139 

140 

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 ] 

150 

151 

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 

156 

157 

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.""" 

164 

165 cursor = await cast("Any", connection).cursor() 

166 try: 

167 _ = await cursor.execute(sql, params) 

168 finally: 

169 await _close_cursor(cursor) 

170 

171 

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] 

184 

185 

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) 

198 

199 

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 

233 

234 

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 

264 

265 

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 ) 

278 

279 

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) 

313 

314 

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 ) 

330 

331 

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.""" 

339 

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