Coverage for tests/test_mariadb_schema.py: 99%

84 statements  

« prev     ^ index     » next       coverage.py v7.14.1, created at 2026-06-01 20:15 +0300

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

2 

3from __future__ import annotations 

4 

5from typing import Any, ClassVar, cast 

6 

7from snektest import assert_eq, assert_raises, load_fixture, test 

8 

9from snekql import ( 

10 MISSING, 

11 Database, 

12 Index, 

13 Pending, 

14 SchemaError, 

15 SchemaVerificationError, 

16 mariadb, 

17) 

18from tests.logging_helpers import NULL_LOGGER 

19from tests.mariadb_server import MariaDBServer, provide_mariadb_server 

20 

21 

22class _RecordingStructuredLogger: 

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

24 

25 def __init__(self) -> None: 

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

27 

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

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

30 

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

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

33 

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

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

36 

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

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

39 

40 

41def _config_from_server(server: MariaDBServer) -> mariadb.Config: 

42 """Build a MariaDB config for the shared local test server.""" 

43 

44 return mariadb.Config( 

45 database=server.database, 

46 host=server.host, 

47 port=server.port, 

48 user=server.user, 

49 ) 

50 

51 

52def _fetch_index_rows( 

53 server: MariaDBServer, table_name: str 

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

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

56 

57 result = server.run_sql( 

58 f""" 

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

60 FROM INFORMATION_SCHEMA.STATISTICS 

61 WHERE TABLE_SCHEMA = DATABASE() 

62 AND TABLE_NAME = '{table_name}' 

63 AND INDEX_NAME <> 'PRIMARY' 

64 GROUP BY INDEX_NAME, NON_UNIQUE 

65 ORDER BY INDEX_NAME 

66 """, 

67 ) 

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

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

70 

71 

72@test(mark="medium") 

73async def mariadb_schema_creates_unique_and_table_indexes() -> None: 

74 """MariaDB startup creates column unique indexes and table indexes.""" 

75 

76 class User[S = Pending](mariadb.Model[S, "User[object]"]): 

77 """Table model with MariaDB indexes.""" 

78 

79 __tablename__ = "issue39_user_indexes" 

80 

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

82 primary_key=True, 

83 auto_increment=True, 

84 default=MISSING, 

85 ) 

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

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

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

89 

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

91 Index(status), 

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

93 ] 

94 

95 server = load_fixture(provide_mariadb_server()) 

96 database = await Database.initialize( 

97 NULL_LOGGER, _config_from_server(server), models=[User] 

98 ) 

99 await database.close() 

100 

101 assert_eq( 

102 _fetch_index_rows(server, "issue39_user_indexes"), 

103 [ 

104 ("ix_issue39_user_indexes_status", "1", "status"), 

105 ("ux_issue39_user_indexes_email", "0", "email"), 

106 ("ux_issue39_user_indexes_tenant_id_email", "0", "tenant_id,email"), 

107 ], 

108 ) 

109 

110 

111@test(mark="medium") 

112async def mariadb_schema_rejects_duplicate_index_names_before_mutation() -> None: 

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

114 

115 class User[S = Pending](mariadb.Model[S, "User[object]"]): 

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

117 

118 __tablename__ = "issue39_duplicate_user" 

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

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

121 

122 class Org[S = Pending](mariadb.Model[S, "Org[object]"]): 

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

124 

125 __tablename__ = "issue39_duplicate_org" 

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

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

128 

129 server = load_fixture(provide_mariadb_server()) 

130 

131 with assert_raises(SchemaError): 

132 _ = await Database.initialize( 

133 NULL_LOGGER, _config_from_server(server), models=[User, Org] 

134 ) 

135 

136 result = server.run_sql( 

137 """ 

138 SELECT COUNT(*) 

139 FROM INFORMATION_SCHEMA.TABLES 

140 WHERE TABLE_SCHEMA = DATABASE() 

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

142 """, 

143 ) 

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

145 

146 

147@test(mark="medium") 

148async def mariadb_strict_schema_policy_raises_on_table_drift() -> None: 

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

150 

151 class User[S = Pending](mariadb.Model[S, "User[object]"]): 

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

153 

154 __tablename__ = "issue39_table_drift" 

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

156 primary_key=True, 

157 auto_increment=True, 

158 default=MISSING, 

159 ) 

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

161 

162 server = load_fixture(provide_mariadb_server()) 

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

164 

165 with assert_raises(SchemaVerificationError): 

166 _ = await Database.initialize( 

167 NULL_LOGGER, _config_from_server(server), models=[User] 

168 ) 

169 

170 

171@test(mark="medium") 

172async def mariadb_strict_schema_policy_raises_on_index_drift() -> None: 

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

174 

175 class User[S = Pending](mariadb.Model[S, "User[object]"]): 

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

177 

178 __tablename__ = "issue39_index_drift" 

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

180 

181 server = load_fixture(provide_mariadb_server()) 

182 _ = server.run_sql( 

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

184 ) 

185 

186 with assert_raises(SchemaVerificationError): 

187 _ = await Database.initialize( 

188 NULL_LOGGER, _config_from_server(server), models=[User] 

189 ) 

190 

191 

192@test(mark="medium") 

193async def mariadb_warn_schema_policy_logs_drift_and_continues() -> None: 

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

195 

196 class User[S = Pending](mariadb.Model[S, "User[object]"]): 

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

198 

199 __tablename__ = "issue39_warn_drift" 

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

201 primary_key=True, 

202 auto_increment=True, 

203 default=MISSING, 

204 ) 

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

206 

207 server = load_fixture(provide_mariadb_server()) 

208 _ = server.run_sql("CREATE TABLE issue39_warn_drift (`email` VARCHAR(255))") 

209 logger = _RecordingStructuredLogger() 

210 database = await Database.initialize( 

211 logger, 

212 _config_from_server(server), 

213 models=[User], 

214 schema_policy="warn", 

215 ) 

216 await database.close() 

217 

218 warnings = [ 

219 fields 

220 for level, event, fields in logger.events 

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

222 ] 

223 assert_eq(len(warnings), 1) 

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