Coverage for tests/mariadb/test_schema.py: 99%

94 statements  

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

1"""MariaDB schema startup and drift verification tests.""" 

2 

3from __future__ import annotations 

4 

5from collections.abc import Sequence 

6from typing import Any, ClassVar, cast 

7 

8from snektest import AsyncFixture, assert_eq, assert_raises, load_fixture, test 

9 

10from snekql import ( 

11 MISSING, 

12 Database, 

13 Fetched, 

14 Index, 

15 Pending, 

16 SchemaError, 

17 SchemaPolicy, 

18 SchemaVerificationError, 

19 StructuredLogger, 

20 mariadb, 

21) 

22from snekql.model import Table 

23from tests.helpers import NULL_LOGGER, TemporaryMariaDBServer, provide_mariadb_server 

24 

25 

26class _RecordingStructuredLogger: 

27 """Structured logger fake that stores event calls for assertions.""" 

28 

29 def __init__(self) -> None: 

30 self.events: list[tuple[str, str, dict[str, object]]] = [] 

31 

32 def debug(self, event: str, **fields: object) -> None: 

33 self.events.append(("debug", event, fields)) 

34 

35 def info(self, event: str, **fields: object) -> None: 

36 self.events.append(("info", event, fields)) 

37 

38 def warning(self, event: str, **fields: object) -> None: 

39 self.events.append(("warning", event, fields)) 

40 

41 def error(self, event: str, **fields: object) -> None: 

42 self.events.append(("error", event, fields)) 

43 

44 

45async def _fetch_index_rows( 

46 server: TemporaryMariaDBServer, table_name: str 

47) -> list[tuple[str, str, str]]: 

48 """Fetch non-primary index metadata from MariaDB information_schema.""" 

49 

50 result = await server.run_sql( 

51 f""" 

52 SELECT INDEX_NAME, NON_UNIQUE, GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) 

53 FROM INFORMATION_SCHEMA.STATISTICS 

54 WHERE TABLE_SCHEMA = DATABASE() 

55 AND TABLE_NAME = '{table_name}' 

56 AND INDEX_NAME <> 'PRIMARY' 

57 GROUP BY INDEX_NAME, NON_UNIQUE 

58 ORDER BY INDEX_NAME 

59 """, 

60 ) 

61 lines = [line for line in result.stdout.splitlines() if line] 

62 return [cast("tuple[str, str, str]", tuple(line.split("\t"))) for line in lines[1:]] 

63 

64 

65async def database_session( 

66 models: Sequence[type[Table[Any]]] = (), 

67 *, 

68 logger: StructuredLogger = NULL_LOGGER, 

69 schema_policy: SchemaPolicy = "strict", 

70 setup_sql: Sequence[str] = (), 

71) -> AsyncFixture[TemporaryMariaDBServer]: 

72 """Provide an initialized MariaDB Database and close it after the test.""" 

73 

74 server = await load_fixture(provide_mariadb_server()) 

75 for sql in setup_sql: 

76 _ = await server.run_sql(sql) 

77 database = await Database.initialize( 

78 server.config(), 

79 logger=logger, 

80 models=models, 

81 schema_policy=schema_policy, 

82 ) 

83 # TODO: do we really need the "try finally" in fixtures? 

84 try: 

85 yield server 

86 finally: 

87 await database.close() 

88 

89 

90@test(mark="medium") 

91async def mariadb_schema_creates_column_unique_indexes() -> None: 

92 """MariaDB startup creates column unique indexes.""" 

93 

94 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]): 

95 """Table model with a MariaDB column unique index.""" 

96 

97 __tablename__ = "issue39_user_column_unique_indexes" 

98 

99 id: User.GenCol[int] = mariadb.Integer( 

100 primary_key=True, 

101 auto_increment=True, 

102 default=MISSING, 

103 ) 

104 email: User.Col[str] = mariadb.Text(nullable=False, unique=True) 

105 

106 server = await load_fixture(database_session([User])) 

107 

108 assert_eq( 

109 await _fetch_index_rows(server, "issue39_user_column_unique_indexes"), 

110 [("ux_issue39_user_column_unique_indexes_email", "0", "email")], 

111 ) 

112 

113 

114@test(mark="medium") 

115async def mariadb_schema_creates_table_indexes() -> None: 

116 """MariaDB startup creates declared table indexes.""" 

117 

118 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]): 

119 """Table model with MariaDB table indexes.""" 

120 

121 __tablename__ = "issue39_user_table_indexes" 

122 

123 id: User.GenCol[int] = mariadb.Integer( 

124 primary_key=True, 

125 auto_increment=True, 

126 default=MISSING, 

127 ) 

128 email: User.Col[str] = mariadb.Text(nullable=False) 

129 status: User.Col[str] = mariadb.Text(nullable=False) 

130 tenant_id: User.Col[int] = mariadb.Integer(nullable=False) 

131 

132 __indexes__: ClassVar[list[Index[Any]]] = [ 

133 Index(status), 

134 Index(tenant_id, email, unique=True), 

135 ] 

136 

137 server = await load_fixture(database_session([User])) 

138 

139 assert_eq( 

140 await _fetch_index_rows(server, "issue39_user_table_indexes"), 

141 [ 

142 ("ix_issue39_user_table_indexes_status", "1", "status"), 

143 ("ux_issue39_user_table_indexes_tenant_id_email", "0", "tenant_id,email"), 

144 ], 

145 ) 

146 

147 

148@test(mark="medium") 

149async def mariadb_schema_rejects_duplicate_index_names_before_mutation() -> None: 

150 """Duplicate resolved index names are rejected before creating tables.""" 

151 

152 server = await load_fixture(provide_mariadb_server()) 

153 

154 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]): 

155 """First model using a duplicate index name.""" 

156 

157 __tablename__ = "issue39_duplicate_user" 

158 email: User.Col[str] = mariadb.Text(nullable=False) 

159 __indexes__: ClassVar[list[Index[Any]]] = [Index(email, name="ix_duplicate")] 

160 

161 class Org[S = Pending](mariadb.Model[S, "Org[Fetched]"]): 

162 """Second model using a duplicate index name.""" 

163 

164 __tablename__ = "issue39_duplicate_org" 

165 name: Org.Col[str] = mariadb.Text(nullable=False) 

166 __indexes__: ClassVar[list[Index[Any]]] = [Index(name, name="ix_duplicate")] 

167 

168 with assert_raises(SchemaError): 

169 _ = await Database.initialize( 

170 server.config(), logger=NULL_LOGGER, models=[User, Org] 

171 ) 

172 

173 result = await server.run_sql( 

174 """ 

175 SELECT COUNT(*) 

176 FROM INFORMATION_SCHEMA.TABLES 

177 WHERE TABLE_SCHEMA = DATABASE() 

178 AND TABLE_NAME IN ('issue39_duplicate_user', 'issue39_duplicate_org') 

179 """, 

180 ) 

181 assert_eq(result.stdout.splitlines()[-1], "0") 

182 

183 

184@test(mark="medium") 

185async def mariadb_strict_schema_policy_raises_on_table_drift() -> None: 

186 """Strict MariaDB schema verification rejects existing table drift.""" 

187 

188 server = await load_fixture(provide_mariadb_server()) 

189 

190 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]): 

191 """Model that expects more columns than the existing table.""" 

192 

193 __tablename__ = "issue39_table_drift" 

194 id: User.GenCol[int] = mariadb.Integer( 

195 primary_key=True, 

196 auto_increment=True, 

197 default=MISSING, 

198 ) 

199 email: User.Col[str] = mariadb.Text(nullable=False) 

200 

201 _ = await server.run_sql("CREATE TABLE issue39_table_drift (`email` VARCHAR(255))") 

202 

203 with assert_raises(SchemaVerificationError): 

204 _ = await Database.initialize( 

205 server.config(), logger=NULL_LOGGER, models=[User] 

206 ) 

207 

208 

209@test(mark="medium") 

210async def mariadb_strict_schema_policy_raises_on_index_drift() -> None: 

211 """Strict MariaDB schema verification rejects missing managed indexes.""" 

212 

213 server = await load_fixture(provide_mariadb_server()) 

214 

215 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]): 

216 """Model that expects a unique index absent from the existing table.""" 

217 

218 __tablename__ = "issue39_index_drift" 

219 email: User.Col[str] = mariadb.Text(nullable=False, unique=True) 

220 

221 _ = await server.run_sql( 

222 "CREATE TABLE issue39_index_drift (`email` VARCHAR(255) NOT NULL)" 

223 ) 

224 

225 with assert_raises(SchemaVerificationError): 

226 _ = await Database.initialize( 

227 server.config(), logger=NULL_LOGGER, models=[User] 

228 ) 

229 

230 

231@test(mark="medium") 

232async def mariadb_warn_schema_policy_logs_drift_and_continues() -> None: 

233 """Warn policy logs MariaDB schema drift without rejecting startup.""" 

234 

235 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]): 

236 """Model used for warn-policy drift verification.""" 

237 

238 __tablename__ = "issue39_warn_drift" 

239 id: User.GenCol[int] = mariadb.Integer( 

240 primary_key=True, 

241 auto_increment=True, 

242 default=MISSING, 

243 ) 

244 email: User.Col[str] = mariadb.Text(nullable=False) 

245 

246 logger = _RecordingStructuredLogger() 

247 _ = await load_fixture( 

248 database_session( 

249 [User], 

250 logger=logger, 

251 schema_policy="warn", 

252 setup_sql=["CREATE TABLE issue39_warn_drift (`email` VARCHAR(255))"], 

253 ) 

254 ) 

255 

256 warnings = [ 

257 fields 

258 for level, event, fields in logger.events 

259 if level == "warning" and event == "schema drift detected" 

260 ] 

261 assert_eq(len(warnings), 1) 

262 assert_eq(warnings[0]["table_name"], "issue39_warn_drift")