Coverage for snekql/runtime.py: 82%

182 statements  

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

1"""Backend-neutral database lifecycle and transaction runtime.""" 

2 

3from __future__ import annotations 

4 

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) 

19 

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 

51 

52SelectOwnerT = TypeVar("SelectOwnerT", bound=Table[Any]) 

53OwnerT = TypeVar("OwnerT", bound=Table[Any]) 

54ReadModelT = TypeVar("ReadModelT", bound=Table[Any]) 

55T = TypeVar("T") 

56Ts = TypeVarTuple("Ts") 

57 

58 

59class RuntimeCursor(Protocol): 

60 """Cursor behavior required by backend-neutral transaction execution.""" 

61 

62 async def fetchone(self) -> Sequence[object] | None: ... 

63 

64 async def fetchall(self) -> Sequence[Sequence[object]]: ... 

65 

66 async def close(self) -> None: ... 

67 

68 

69class RuntimeConnection(Protocol): 

70 """Connection behavior required by backend-neutral transactions.""" 

71 

72 async def begin(self) -> None: ... 

73 

74 async def commit(self) -> None: ... 

75 

76 async def rollback(self) -> None: ... 

77 

78 async def execute( 

79 self, 

80 sql: str, 

81 params: tuple[object, ...], 

82 ) -> RuntimeCursor: ... 

83 

84 

85class RuntimeBackend(Protocol): 

86 """Backend adapter seam used by Database and Transaction.""" 

87 

88 acquire_timeout: NonNegativeFloat 

89 backend_family: BackendFamily 

90 logger: ResolvedStructuredLogger 

91 

92 async def acquire( 

93 self, 

94 acquisition_timeout: NonNegativeFloat, 

95 ) -> RuntimeConnection: ... 

96 

97 async def release(self, connection: object) -> None: ... 

98 

99 async def close(self, close_timeout: NonNegativeFloat) -> None: ... 

100 

101 def check_accepting_work(self) -> None: ... 

102 

103 def compile_select_sql( 

104 self, 

105 query: AnySelectQuery, 

106 ) -> tuple[str, tuple[object, ...]]: ... 

107 

108 def compile_write_sql(self, query: object) -> tuple[str, tuple[object, ...]]: ... 

109 

110 def materialize_select_row( 

111 self, 

112 query: AnySelectQuery, 

113 row: Sequence[object], 

114 ) -> object: ... 

115 

116 

117class Transaction: 

118 """Async transaction that executes built snekql queries on one connection. 

119 

120 >>> async def create_user(transaction: Transaction, user: User[Pending]) -> None: 

121 ... await transaction.execute(insert(user)) 

122 """ 

123 

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 

137 

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 

165 

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 ) 

209 

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.""" 

222 

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 ] 

256 

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.""" 

269 

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)) 

302 

303 async def execute( 

304 self, query: InsertQuery[Any] | UpdateQuery[Any] | DeleteQuery[Any] 

305 ) -> None: 

306 """Execute a write query inside this transaction.""" 

307 

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 ) 

335 

336 def require_connection(self) -> RuntimeConnection: 

337 """Return the active transaction connection or reject use-after-close.""" 

338 

339 connection = self.connection 

340 if self.closed or connection is None: 

341 msg = "transaction is closed" 

342 raise TransactionClosedError(msg) 

343 return connection 

344 

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) 

356 

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) 

368 

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) 

375 

376 

377class Database: 

378 """Initialized snekql runtime service for database-backed execution. 

379 

380 `Database.initialize(..., logger=logger)` is the only public construction 

381 path. A Database owns connectivity, schema startup work, and transaction entry. 

382 """ 

383 

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) 

388 

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: ... 

399 

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: ... 

410 

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: ... 

423 

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.""" 

437 

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 

486 

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.""" 

490 

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 ) 

499 

500 async def close(self) -> None: 

501 """Close this database runtime idempotently when shutdown succeeds.""" 

502 

503 await self.runtime.close(self.runtime.acquire_timeout)