Coverage for snekql/runtime.py: 82%
182 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"""Backend-neutral database lifecycle and transaction runtime."""
3from __future__ import annotations
5from collections.abc import Sequence
6from pathlib import Path
7from types import TracebackType
8from typing import (
9 Any,
10 Literal,
11 Never,
12 Protocol,
13 Self,
14 TypeVar,
15 TypeVarTuple,
16 cast,
17 overload,
18)
20from snekql._runtime_selection import resolve_runtime_selection
21from snekql.errors import (
22 DatabaseRuntimeError,
23 ExecutionError,
24 QueryCompilationError,
25 TransactionClosedError,
26)
27from snekql.mariadb.config import Config as MariaDBConfig
28from snekql.model import (
29 BackendFamily,
30 Table,
31 require_model_backend,
32 require_model_table_name,
33)
34from snekql.query import (
35 AnySelectQuery,
36 DeleteQuery,
37 InsertQuery,
38 SelectModelQuery,
39 SelectTupleQuery,
40 SelectValueQuery,
41 UpdateQuery,
42)
43from snekql.sqlite.config import Config as SQLiteConfig
44from snekql.storage import SchemaPolicy
45from snekql.structured_logging import (
46 ResolvedStructuredLogger,
47 StructuredLogger,
48 resolve_structured_logger,
49)
50from snekql.validation import NonNegativeFloat, PositiveInt, validate_boundary
52SelectOwnerT = TypeVar("SelectOwnerT", bound=Table[Any])
53OwnerT = TypeVar("OwnerT", bound=Table[Any])
54ReadModelT = TypeVar("ReadModelT", bound=Table[Any])
55T = TypeVar("T")
56Ts = TypeVarTuple("Ts")
59class RuntimeCursor(Protocol):
60 """Cursor behavior required by backend-neutral transaction execution."""
62 async def fetchone(self) -> Sequence[object] | None: ...
64 async def fetchall(self) -> Sequence[Sequence[object]]: ...
66 async def close(self) -> None: ...
69class RuntimeConnection(Protocol):
70 """Connection behavior required by backend-neutral transactions."""
72 async def begin(self) -> None: ...
74 async def commit(self) -> None: ...
76 async def rollback(self) -> None: ...
78 async def execute(
79 self,
80 sql: str,
81 params: tuple[object, ...],
82 ) -> RuntimeCursor: ...
85class RuntimeBackend(Protocol):
86 """Backend adapter seam used by Database and Transaction."""
88 acquire_timeout: NonNegativeFloat
89 backend_family: BackendFamily
90 logger: ResolvedStructuredLogger
92 async def acquire(
93 self,
94 acquisition_timeout: NonNegativeFloat,
95 ) -> RuntimeConnection: ...
97 async def release(self, connection: object) -> None: ...
99 async def close(self, close_timeout: NonNegativeFloat) -> None: ...
101 def check_accepting_work(self) -> None: ...
103 def compile_select_sql(
104 self,
105 query: AnySelectQuery,
106 ) -> tuple[str, tuple[object, ...]]: ...
108 def compile_write_sql(self, query: object) -> tuple[str, tuple[object, ...]]: ...
110 def materialize_select_row(
111 self,
112 query: AnySelectQuery,
113 row: Sequence[object],
114 ) -> object: ...
117class Transaction:
118 """Async transaction that executes built snekql queries on one connection.
120 >>> async def create_user(transaction: Transaction, user: User[Pending]) -> None:
121 ... await transaction.execute(insert(user))
122 """
124 def __init__(
125 self,
126 *,
127 runtime: RuntimeBackend | None = None,
128 timeout: NonNegativeFloat = 0.0,
129 ) -> None:
130 if runtime is None:
131 msg = "use db.transaction(...) to start a transaction"
132 raise DatabaseRuntimeError(msg)
133 self.closed: bool = False
134 self.connection: RuntimeConnection | None = None
135 self.runtime: RuntimeBackend = runtime
136 self.timeout: NonNegativeFloat = timeout
138 async def __aenter__(self) -> Self:
139 if self.closed or self.connection is not None:
140 msg = "transaction is closed"
141 raise TransactionClosedError(msg)
142 self.runtime.logger.debug(
143 "transaction acquiring connection",
144 backend=self.runtime.backend_family,
145 timeout=self.timeout,
146 )
147 connection = await self.runtime.acquire(self.timeout)
148 try:
149 await connection.begin()
150 except Exception as error:
151 self.runtime.logger.error( # noqa: TRY400
152 "transaction begin failed",
153 backend=self.runtime.backend_family,
154 error_type=type(error).__name__,
155 )
156 await self.runtime.release(connection)
157 msg = "could not begin transaction"
158 raise DatabaseRuntimeError(msg) from error
159 self.connection = connection
160 self.runtime.logger.debug(
161 "transaction begin",
162 backend=self.runtime.backend_family,
163 )
164 return self
166 async def __aexit__(
167 self,
168 exc_type: type[BaseException] | None,
169 exc_value: BaseException | None,
170 traceback: TracebackType | None,
171 ) -> None:
172 _ = exc_value
173 _ = traceback
174 connection = self.connection
175 if connection is None:
176 msg = "transaction is closed"
177 raise TransactionClosedError(msg)
178 self.connection = None
179 self.closed = True
180 try:
181 if exc_type is None:
182 await connection.commit()
183 self.runtime.logger.debug(
184 "transaction commit",
185 backend=self.runtime.backend_family,
186 )
187 else:
188 await connection.rollback()
189 self.runtime.logger.debug(
190 "transaction rollback",
191 backend=self.runtime.backend_family,
192 exception_type=exc_type.__name__,
193 )
194 except Exception as error:
195 self.runtime.logger.error( # noqa: TRY400
196 "transaction close failed",
197 backend=self.runtime.backend_family,
198 error_type=type(error).__name__,
199 )
200 if exc_type is None:
201 msg = "could not close transaction"
202 raise DatabaseRuntimeError(msg) from error
203 finally:
204 await self.runtime.release(connection)
205 self.runtime.logger.debug(
206 "transaction released",
207 backend=self.runtime.backend_family,
208 )
210 @overload
211 async def fetch_all(
212 self, query: SelectModelQuery[SelectOwnerT, ReadModelT]
213 ) -> list[ReadModelT]: ...
214 @overload
215 async def fetch_all(self, query: SelectValueQuery[OwnerT, T]) -> list[T]: ...
216 @overload
217 async def fetch_all(
218 self, query: SelectTupleQuery[OwnerT, *Ts]
219 ) -> list[tuple[*Ts]]: ...
220 async def fetch_all(self, query: object) -> object:
221 """Fetch all rows for a select query."""
223 connection = self.require_connection()
224 select_query = self._require_select_query(query)
225 self._validate_query_backend(select_query)
226 sql, params = self.runtime.compile_select_sql(select_query)
227 try:
228 cursor = await connection.execute(sql, params)
229 try:
230 rows = await cursor.fetchall()
231 finally:
232 await cursor.close()
233 except Exception as error:
234 self.runtime.logger.error( # noqa: TRY400
235 "query failed",
236 backend=self.runtime.backend_family,
237 error_type=type(error).__name__,
238 operation="fetch_all",
239 params=params,
240 sql=sql,
241 )
242 msg = "select failed"
243 raise ExecutionError(msg, sql=sql, params=params) from error
244 self.runtime.logger.debug(
245 "query executed",
246 backend=self.runtime.backend_family,
247 operation="fetch_all",
248 params=params,
249 row_count=len(rows),
250 sql=sql,
251 )
252 return [
253 self.runtime.materialize_select_row(select_query, tuple(row))
254 for row in rows
255 ]
257 @overload
258 async def fetch_one(
259 self, query: SelectModelQuery[SelectOwnerT, ReadModelT]
260 ) -> ReadModelT | None: ...
261 @overload
262 async def fetch_one(self, query: SelectValueQuery[OwnerT, T]) -> T | None: ...
263 @overload
264 async def fetch_one(
265 self, query: SelectTupleQuery[OwnerT, *Ts]
266 ) -> tuple[*Ts] | None: ...
267 async def fetch_one(self, query: object) -> object:
268 """Fetch one row for a select query."""
270 connection = self.require_connection()
271 select_query = self._require_select_query(query)
272 self._validate_query_backend(select_query)
273 sql, params = self.runtime.compile_select_sql(select_query)
274 try:
275 cursor = await connection.execute(sql, params)
276 try:
277 row = await cursor.fetchone()
278 finally:
279 await cursor.close()
280 except Exception as error:
281 self.runtime.logger.error( # noqa: TRY400
282 "query failed",
283 backend=self.runtime.backend_family,
284 error_type=type(error).__name__,
285 operation="fetch_one",
286 params=params,
287 sql=sql,
288 )
289 msg = "select failed"
290 raise ExecutionError(msg, sql=sql, params=params) from error
291 self.runtime.logger.debug(
292 "query executed",
293 backend=self.runtime.backend_family,
294 operation="fetch_one",
295 params=params,
296 row_found=row is not None,
297 sql=sql,
298 )
299 if row is None:
300 return None
301 return self.runtime.materialize_select_row(select_query, tuple(row))
303 async def execute(
304 self, query: InsertQuery[Any] | UpdateQuery[Any] | DeleteQuery[Any]
305 ) -> None:
306 """Execute a write query inside this transaction."""
308 connection = self.require_connection()
309 self._validate_query_backend(query)
310 sql, params = self.runtime.compile_write_sql(query)
311 try:
312 cursor = await connection.execute(sql, params)
313 try:
314 pass
315 finally:
316 await cursor.close()
317 except Exception as error:
318 self.runtime.logger.error( # noqa: TRY400
319 "query failed",
320 backend=self.runtime.backend_family,
321 error_type=type(error).__name__,
322 operation="write",
323 params=params,
324 sql=sql,
325 )
326 msg = "write failed"
327 raise ExecutionError(msg, sql=sql, params=params) from error
328 self.runtime.logger.debug(
329 "query executed",
330 backend=self.runtime.backend_family,
331 operation="write",
332 params=params,
333 sql=sql,
334 )
336 def require_connection(self) -> RuntimeConnection:
337 """Return the active transaction connection or reject use-after-close."""
339 connection = self.connection
340 if self.closed or connection is None:
341 msg = "transaction is closed"
342 raise TransactionClosedError(msg)
343 return connection
345 def _validate_query_backend(self, query: object) -> None:
346 query_model = self._query_model(query)
347 received_backend = require_model_backend(query_model)
348 expected_backend = self.runtime.backend_family
349 if received_backend == expected_backend:
350 return
351 msg = (
352 f"backend mismatch: expected {expected_backend} query, "
353 f"received {received_backend} query for {query_model.__name__}"
354 )
355 raise DatabaseRuntimeError(msg)
357 @staticmethod
358 def _query_model(query: object) -> type[Table[Any]]:
359 if isinstance(query, InsertQuery):
360 insert_query = cast("InsertQuery[Any]", query)
361 return cast("type[Table[Any]]", type(insert_query.row))
362 if isinstance(query, SelectModelQuery | SelectValueQuery | SelectTupleQuery):
363 return query.state.model
364 if isinstance(query, UpdateQuery | DeleteQuery):
365 return query.state.model
366 msg = "query backend validation requires a snekql query"
367 raise QueryCompilationError(msg)
369 @staticmethod
370 def _require_select_query(query: object) -> AnySelectQuery:
371 if isinstance(query, SelectModelQuery | SelectValueQuery | SelectTupleQuery):
372 return cast("AnySelectQuery", query)
373 msg = "fetch requires a select query"
374 raise QueryCompilationError(msg)
377class Database:
378 """Initialized snekql runtime service for database-backed execution.
380 `Database.initialize(..., logger=logger)` is the only public construction
381 path. A Database owns connectivity, schema startup work, and transaction entry.
382 """
384 def __init__(self, _initialized: Never, /) -> None:
385 self.runtime = cast("RuntimeBackend", None)
386 msg = "use Database.initialize(..., logger=logger) to create a Database"
387 raise DatabaseRuntimeError(msg)
389 @overload
390 @classmethod
391 async def initialize(
392 cls,
393 backend: SQLiteConfig,
394 *,
395 logger: StructuredLogger,
396 models: Sequence[type[Table[Any]]] = (),
397 schema_policy: SchemaPolicy = "strict",
398 ) -> Self: ...
400 @overload
401 @classmethod
402 async def initialize(
403 cls,
404 backend: MariaDBConfig,
405 *,
406 logger: StructuredLogger,
407 models: Sequence[type[Table[Any]]] = (),
408 schema_policy: SchemaPolicy = "strict",
409 ) -> Self: ...
411 @overload
412 @classmethod
413 async def initialize(
414 cls,
415 *,
416 logger: StructuredLogger,
417 database: Path | Literal[":memory:"],
418 models: Sequence[type[Table[Any]]] = (),
419 schema_policy: SchemaPolicy = "strict",
420 pool_size: PositiveInt = 5,
421 acquire_timeout: NonNegativeFloat = 30.0,
422 ) -> Self: ...
424 @classmethod
425 async def initialize( # noqa: PLR0913
426 cls,
427 backend: object | None = None,
428 *,
429 logger: StructuredLogger,
430 database: Path | Literal[":memory:"] | None = None,
431 models: Sequence[type[Table[Any]]] = (),
432 schema_policy: SchemaPolicy = "strict",
433 pool_size: PositiveInt = 5,
434 acquire_timeout: NonNegativeFloat = 30.0,
435 ) -> Self:
436 """Initialize connectivity, schema startup, and runtime lifecycle."""
438 structured_logger = resolve_structured_logger(logger=logger)
439 try:
440 runtime_selection = resolve_runtime_selection(
441 backend=backend,
442 database=database,
443 pool_size=pool_size,
444 acquire_timeout=acquire_timeout,
445 )
446 runtime_config = runtime_selection.config
447 backend_family = runtime_selection.backend_family
448 runtime_selection.validate_model_backends(models)
449 table_names = tuple(require_model_table_name(model) for model in models)
450 structured_logger.info(
451 "database initialization started",
452 backend=backend_family,
453 model_count=len(models),
454 schema_policy=schema_policy,
455 table_names=table_names,
456 )
457 structured_logger.debug(
458 "database backend selected",
459 backend=backend_family,
460 acquire_timeout=runtime_config.acquire_timeout,
461 pool_size=runtime_config.pool_size,
462 )
463 runtime = cast(
464 "RuntimeBackend",
465 await runtime_selection.initialize_runtime(
466 models,
467 schema_policy,
468 logger=structured_logger,
469 ),
470 )
471 structured_logger.info(
472 "database initialization completed",
473 backend=backend_family,
474 model_count=len(models),
475 table_names=table_names,
476 )
477 except Exception as error:
478 structured_logger.error( # noqa: TRY400
479 "database initialization failed",
480 error_type=type(error).__name__,
481 )
482 raise
483 database_instance = cls.__new__(cls)
484 database_instance.runtime = runtime
485 return database_instance
487 @validate_boundary(error_type=DatabaseRuntimeError)
488 def transaction(self, *, timeout: NonNegativeFloat | None = None) -> Transaction:
489 """Create a transaction context manager using the runtime backend."""
491 self.runtime.check_accepting_work()
492 acquisition_timeout = (
493 self.runtime.acquire_timeout if timeout is None else timeout
494 )
495 return Transaction(
496 runtime=self.runtime,
497 timeout=acquisition_timeout,
498 )
500 async def close(self) -> None:
501 """Close this database runtime idempotently when shutdown succeeds."""
503 await self.runtime.close(self.runtime.acquire_timeout)