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
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-01 20:15 +0300
1"""MariaDB schema startup and drift verification tests."""
3from __future__ import annotations
5from typing import Any, ClassVar, cast
7from snektest import assert_eq, assert_raises, load_fixture, test
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
22class _RecordingStructuredLogger:
23 """Structured logger fake that stores event calls for assertions."""
25 def __init__(self) -> None:
26 self.events: list[tuple[str, str, dict[str, object]]] = []
28 def debug(self, event: str, **fields: object) -> None:
29 self.events.append(("debug", event, fields))
31 def info(self, event: str, **fields: object) -> None:
32 self.events.append(("info", event, fields))
34 def warning(self, event: str, **fields: object) -> None:
35 self.events.append(("warning", event, fields))
37 def error(self, event: str, **fields: object) -> None:
38 self.events.append(("error", event, fields))
41def _config_from_server(server: MariaDBServer) -> mariadb.Config:
42 """Build a MariaDB config for the shared local test server."""
44 return mariadb.Config(
45 database=server.database,
46 host=server.host,
47 port=server.port,
48 user=server.user,
49 )
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."""
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:]]
72@test(mark="medium")
73async def mariadb_schema_creates_unique_and_table_indexes() -> None:
74 """MariaDB startup creates column unique indexes and table indexes."""
76 class User[S = Pending](mariadb.Model[S, "User[object]"]):
77 """Table model with MariaDB indexes."""
79 __tablename__ = "issue39_user_indexes"
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)
90 __indexes__: ClassVar[list[Index[Any]]] = [
91 Index(status),
92 Index(tenant_id, email, unique=True),
93 ]
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()
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 )
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."""
115 class User[S = Pending](mariadb.Model[S, "User[object]"]):
116 """First model using a duplicate index name."""
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")]
122 class Org[S = Pending](mariadb.Model[S, "Org[object]"]):
123 """Second model using a duplicate index name."""
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")]
129 server = load_fixture(provide_mariadb_server())
131 with assert_raises(SchemaError):
132 _ = await Database.initialize(
133 NULL_LOGGER, _config_from_server(server), models=[User, Org]
134 )
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")
147@test(mark="medium")
148async def mariadb_strict_schema_policy_raises_on_table_drift() -> None:
149 """Strict MariaDB schema verification rejects existing table drift."""
151 class User[S = Pending](mariadb.Model[S, "User[object]"]):
152 """Model that expects more columns than the existing table."""
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)
162 server = load_fixture(provide_mariadb_server())
163 _ = server.run_sql("CREATE TABLE issue39_table_drift (`email` VARCHAR(255))")
165 with assert_raises(SchemaVerificationError):
166 _ = await Database.initialize(
167 NULL_LOGGER, _config_from_server(server), models=[User]
168 )
171@test(mark="medium")
172async def mariadb_strict_schema_policy_raises_on_index_drift() -> None:
173 """Strict MariaDB schema verification rejects missing managed indexes."""
175 class User[S = Pending](mariadb.Model[S, "User[object]"]):
176 """Model that expects a unique index absent from the existing table."""
178 __tablename__ = "issue39_index_drift"
179 email: User.Col[str] = mariadb.Text(nullable=False, unique=True)
181 server = load_fixture(provide_mariadb_server())
182 _ = server.run_sql(
183 "CREATE TABLE issue39_index_drift (`email` VARCHAR(255) NOT NULL)"
184 )
186 with assert_raises(SchemaVerificationError):
187 _ = await Database.initialize(
188 NULL_LOGGER, _config_from_server(server), models=[User]
189 )
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."""
196 class User[S = Pending](mariadb.Model[S, "User[object]"]):
197 """Model used for warn-policy drift verification."""
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)
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()
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")