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
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-07 21:13 +0300
1"""MariaDB schema startup and drift verification tests."""
3from __future__ import annotations
5from collections.abc import Sequence
6from typing import Any, ClassVar, cast
8from snektest import AsyncFixture, assert_eq, assert_raises, load_fixture, test
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
26class _RecordingStructuredLogger:
27 """Structured logger fake that stores event calls for assertions."""
29 def __init__(self) -> None:
30 self.events: list[tuple[str, str, dict[str, object]]] = []
32 def debug(self, event: str, **fields: object) -> None:
33 self.events.append(("debug", event, fields))
35 def info(self, event: str, **fields: object) -> None:
36 self.events.append(("info", event, fields))
38 def warning(self, event: str, **fields: object) -> None:
39 self.events.append(("warning", event, fields))
41 def error(self, event: str, **fields: object) -> None:
42 self.events.append(("error", event, fields))
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."""
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:]]
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."""
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()
90@test(mark="medium")
91async def mariadb_schema_creates_column_unique_indexes() -> None:
92 """MariaDB startup creates column unique indexes."""
94 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]):
95 """Table model with a MariaDB column unique index."""
97 __tablename__ = "issue39_user_column_unique_indexes"
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)
106 server = await load_fixture(database_session([User]))
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 )
114@test(mark="medium")
115async def mariadb_schema_creates_table_indexes() -> None:
116 """MariaDB startup creates declared table indexes."""
118 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]):
119 """Table model with MariaDB table indexes."""
121 __tablename__ = "issue39_user_table_indexes"
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)
132 __indexes__: ClassVar[list[Index[Any]]] = [
133 Index(status),
134 Index(tenant_id, email, unique=True),
135 ]
137 server = await load_fixture(database_session([User]))
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 )
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."""
152 server = await load_fixture(provide_mariadb_server())
154 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]):
155 """First model using a duplicate index name."""
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")]
161 class Org[S = Pending](mariadb.Model[S, "Org[Fetched]"]):
162 """Second model using a duplicate index name."""
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")]
168 with assert_raises(SchemaError):
169 _ = await Database.initialize(
170 server.config(), logger=NULL_LOGGER, models=[User, Org]
171 )
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")
184@test(mark="medium")
185async def mariadb_strict_schema_policy_raises_on_table_drift() -> None:
186 """Strict MariaDB schema verification rejects existing table drift."""
188 server = await load_fixture(provide_mariadb_server())
190 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]):
191 """Model that expects more columns than the existing table."""
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)
201 _ = await server.run_sql("CREATE TABLE issue39_table_drift (`email` VARCHAR(255))")
203 with assert_raises(SchemaVerificationError):
204 _ = await Database.initialize(
205 server.config(), logger=NULL_LOGGER, models=[User]
206 )
209@test(mark="medium")
210async def mariadb_strict_schema_policy_raises_on_index_drift() -> None:
211 """Strict MariaDB schema verification rejects missing managed indexes."""
213 server = await load_fixture(provide_mariadb_server())
215 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]):
216 """Model that expects a unique index absent from the existing table."""
218 __tablename__ = "issue39_index_drift"
219 email: User.Col[str] = mariadb.Text(nullable=False, unique=True)
221 _ = await server.run_sql(
222 "CREATE TABLE issue39_index_drift (`email` VARCHAR(255) NOT NULL)"
223 )
225 with assert_raises(SchemaVerificationError):
226 _ = await Database.initialize(
227 server.config(), logger=NULL_LOGGER, models=[User]
228 )
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."""
235 class User[S = Pending](mariadb.Model[S, "User[Fetched]"]):
236 """Model used for warn-policy drift verification."""
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)
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 )
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")