Coverage for snekql/schema.py: 95%

114 statements  

« prev     ^ index     » next       coverage.py v7.14.1, created at 2026-06-07 21:13 +0300

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

2 

3from __future__ import annotations 

4 

5import contextlib 

6from collections.abc import Sequence 

7from typing import Any 

8 

9from aiosqlite import Connection, Error 

10 

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 

25 

26 

27def quote_sqlite_identifier(identifier: str) -> str: 

28 """Quote a SQLite identifier with double-quote escaping.""" 

29 

30 return quote_identifier(identifier) 

31 

32 

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) 

47 

48 

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

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

51 

52 

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 ) 

62 

63 

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 ) 

73 

74 

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 ] 

80 

81 

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] 

99 

100 

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 

122 

123 

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(" ,", ",") 

128 

129 

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 ) 

142 

143 

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) 

159 

160 

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) 

192 

193 

194async def _rollback_schema_setup(connection: Connection) -> None: 

195 with contextlib.suppress(Error): 

196 _ = await connection.execute("ROLLBACK") 

197 

198 

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 

228 

229 

230def validate_schema_models(models: Sequence[type[Table[Any]]]) -> None: 

231 """Reject duplicate resolved table names before schema startup.""" 

232 

233 _ = build_schema_plan(models) 

234 

235 

236def validate_schema_policy(schema_policy: SchemaPolicy) -> None: 

237 """Reject unsupported schema policy values.""" 

238 

239 validate_planned_schema_policy(schema_policy) 

240 

241 

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

249 

250 await _initialize_sqlite_schema( 

251 connection, 

252 models, 

253 schema_policy, 

254 logger, 

255 )