Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/database/configurations.py: 55%

240 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 02:04 +0000

1from __future__ import annotations 

2 

3import logging 

4import sqlite3 

5import ssl 

6import traceback 

7from abc import ABC, abstractmethod 

8from asyncio import AbstractEventLoop, get_running_loop 

9from collections.abc import AsyncGenerator, Hashable 

10from contextlib import AbstractAsyncContextManager, asynccontextmanager 

11from contextvars import ContextVar 

12from functools import partial 

13from typing import Any, Optional 

14 

15import sqlalchemy as sa 

16from sqlalchemy import AdaptedConnection, event 

17from sqlalchemy.dialects.sqlite import aiosqlite 

18from sqlalchemy.engine.interfaces import DBAPIConnection 

19from sqlalchemy.ext.asyncio import ( 

20 AsyncConnection, 

21 AsyncEngine, 

22 AsyncSession, 

23 AsyncSessionTransaction, 

24 create_async_engine, 

25) 

26from sqlalchemy.pool import ConnectionPoolEntry 

27from typing_extensions import TypeAlias 

28 

29from prefect._internal.observability import configure_logfire 

30from prefect._internal.plugins.manager import ( 

31 build_manager, 

32 call_async_hook, 

33 load_entry_point_plugins, 

34) 

35from prefect.plugins import HookSpec 

36from prefect.settings import ( 

37 PREFECT_API_DATABASE_CONNECTION_TIMEOUT, 

38 PREFECT_API_DATABASE_ECHO, 

39 PREFECT_API_DATABASE_TIMEOUT, 

40 PREFECT_TESTING_UNIT_TEST_MODE, 

41 get_current_settings, 

42) 

43from prefect.utilities.asyncutils import add_event_loop_shutdown_callback 

44 

45logfire: Any | None = configure_logfire() 

46 

47SQLITE_BEGIN_MODE: ContextVar[Optional[str]] = ContextVar( # novm 

48 "SQLITE_BEGIN_MODE", default=None 

49) 

50 

51_EngineCacheKey: TypeAlias = tuple[AbstractEventLoop, str, bool, Optional[float]] 

52ENGINES: dict[_EngineCacheKey, AsyncEngine] = {} 

53 

54 

55class ConnectionTracker: 

56 """A test utility which tracks the connections given out by a connection pool, to 

57 make it easy to see which connections are currently checked out and open.""" 

58 

59 all_connections: dict[AdaptedConnection, list[str]] 

60 open_connections: dict[AdaptedConnection, list[str]] 

61 left_field_closes: dict[AdaptedConnection, list[str]] 

62 connects: int 

63 closes: int 

64 active: bool 

65 

66 def __init__(self) -> None: 

67 self.active = False 

68 self.all_connections = {} 

69 self.open_connections = {} 

70 self.left_field_closes = {} 

71 self.connects = 0 

72 self.closes = 0 

73 

74 def track_pool(self, pool: sa.pool.Pool) -> None: 

75 event.listen(pool, "connect", self.on_connect) 

76 event.listen(pool, "close", self.on_close) 

77 event.listen(pool, "close_detached", self.on_close_detached) 

78 

79 def on_connect( 

80 self, 

81 adapted_connection: AdaptedConnection, 

82 connection_record: ConnectionPoolEntry, 

83 ) -> None: 

84 self.all_connections[adapted_connection] = traceback.format_stack() 

85 self.open_connections[adapted_connection] = traceback.format_stack() 

86 self.connects += 1 

87 

88 def on_close( 

89 self, 

90 adapted_connection: AdaptedConnection, 

91 connection_record: ConnectionPoolEntry, 

92 ) -> None: 

93 try: 

94 del self.open_connections[adapted_connection] 

95 except KeyError: 

96 self.left_field_closes[adapted_connection] = traceback.format_stack() 

97 self.closes += 1 

98 

99 def on_close_detached( 

100 self, 

101 adapted_connection: AdaptedConnection, 

102 ) -> None: 

103 try: 

104 del self.open_connections[adapted_connection] 

105 except KeyError: 

106 self.left_field_closes[adapted_connection] = traceback.format_stack() 

107 self.closes += 1 

108 

109 def clear(self) -> None: 

110 self.all_connections.clear() 

111 self.open_connections.clear() 

112 self.left_field_closes.clear() 

113 self.connects = 0 

114 self.closes = 0 

115 

116 

117TRACKER: ConnectionTracker = ConnectionTracker() 

118 

119 

120class BaseDatabaseConfiguration(ABC): 

121 """ 

122 Abstract base class used to inject database connection configuration into Prefect. 

123 

124 This configuration is responsible for defining how Prefect REST API creates and manages 

125 database connections and sessions. 

126 """ 

127 

128 def __init__( 

129 self, 

130 connection_url: str, 

131 echo: Optional[bool] = None, 

132 timeout: Optional[float] = None, 

133 connection_timeout: Optional[float] = None, 

134 sqlalchemy_pool_size: Optional[int] = None, 

135 sqlalchemy_max_overflow: Optional[int] = None, 

136 connection_app_name: Optional[str] = None, 

137 statement_cache_size: Optional[int] = None, 

138 prepared_statement_cache_size: Optional[int] = None, 

139 search_path: Optional[str] = None, 

140 ) -> None: 

141 self.connection_url = connection_url 

142 self.echo: bool = echo or PREFECT_API_DATABASE_ECHO.value() 

143 self.timeout: Optional[float] = timeout or PREFECT_API_DATABASE_TIMEOUT.value() 

144 self.connection_timeout: Optional[float] = ( 

145 connection_timeout or PREFECT_API_DATABASE_CONNECTION_TIMEOUT.value() 

146 ) 

147 self.sqlalchemy_pool_size: Optional[int] = ( 

148 sqlalchemy_pool_size 

149 or get_current_settings().server.database.sqlalchemy.pool_size 

150 ) 

151 self.sqlalchemy_max_overflow: Optional[int] = ( 

152 sqlalchemy_max_overflow 

153 or get_current_settings().server.database.sqlalchemy.max_overflow 

154 ) 

155 self.connection_app_name: Optional[str] = ( 

156 connection_app_name 

157 or get_current_settings().server.database.sqlalchemy.connect_args.application_name 

158 ) 

159 self.statement_cache_size: Optional[int] = ( 

160 statement_cache_size 

161 or get_current_settings().server.database.sqlalchemy.connect_args.statement_cache_size 

162 ) 

163 self.prepared_statement_cache_size: Optional[int] = ( 

164 prepared_statement_cache_size 

165 or get_current_settings().server.database.sqlalchemy.connect_args.prepared_statement_cache_size 

166 ) 

167 self.search_path: Optional[str] = ( 

168 search_path 

169 or get_current_settings().server.database.sqlalchemy.connect_args.search_path 

170 ) 

171 

172 def unique_key(self) -> tuple[Hashable, ...]: 

173 """ 

174 Returns a key used to determine whether to instantiate a new DB interface. 

175 """ 

176 return (self.__class__, self.connection_url) 

177 

178 @abstractmethod 

179 async def engine(self) -> AsyncEngine: 

180 """Returns a SqlAlchemy engine""" 

181 

182 @abstractmethod 

183 async def session(self, engine: AsyncEngine) -> AsyncSession: 

184 """ 

185 Retrieves a SQLAlchemy session for an engine. 

186 """ 

187 

188 @abstractmethod 

189 async def create_db( 

190 self, connection: AsyncConnection, base_metadata: sa.MetaData 

191 ) -> None: 

192 """Create the database""" 

193 

194 @abstractmethod 

195 async def drop_db( 

196 self, connection: AsyncConnection, base_metadata: sa.MetaData 

197 ) -> None: 

198 """Drop the database""" 

199 

200 @abstractmethod 

201 def is_inmemory(self) -> bool: 

202 """Returns true if database is run in memory""" 

203 

204 @abstractmethod 

205 def begin_transaction( 

206 self, session: AsyncSession, with_for_update: bool = False 

207 ) -> AbstractAsyncContextManager[AsyncSessionTransaction]: 

208 """Enter a transaction for a session""" 

209 pass 

210 

211 

212class AsyncPostgresConfiguration(BaseDatabaseConfiguration): 

213 async def engine(self) -> AsyncEngine: 

214 """Retrieves an async SQLAlchemy engine. 

215 

216 Args: 

217 connection_url (str, optional): The database connection string. 

218 Defaults to self.connection_url 

219 echo (bool, optional): Whether to echo SQL sent 

220 to the database. Defaults to self.echo 

221 timeout (float, optional): The database statement timeout, in seconds. 

222 Defaults to self.timeout 

223 

224 Returns: 

225 AsyncEngine: a SQLAlchemy engine 

226 """ 

227 

228 loop = get_running_loop() 

229 

230 cache_key = ( 

231 loop, 

232 self.connection_url, 

233 self.echo, 

234 self.timeout, 

235 ) 

236 if cache_key not in ENGINES: 

237 kwargs: dict[str, Any] = ( 

238 get_current_settings().server.database.sqlalchemy.model_dump( 

239 mode="json", exclude={"connect_args"} 

240 ) 

241 ) 

242 connect_args: dict[str, Any] = {} 

243 

244 if self.timeout is not None: 

245 connect_args["command_timeout"] = self.timeout 

246 

247 # In test mode, use a higher connection timeout to handle the heavy 

248 # load of parallel test execution (pytest-xdist). Establishing a new 

249 # asyncpg connection can occasionally take longer than the 5s default 

250 # under CI load, which surfaces as a TimeoutError during fixture 

251 # setup. Keep the configured value if the user has already raised it. 

252 connection_timeout = self.connection_timeout 

253 if PREFECT_TESTING_UNIT_TEST_MODE.value() is True: 253 ↛ 254line 253 didn't jump to line 254 because the condition on line 253 was never true

254 connection_timeout = max(connection_timeout or 0.0, 30.0) 

255 

256 if connection_timeout is not None: 256 ↛ 259line 256 didn't jump to line 259 because the condition on line 256 was always true

257 connect_args["timeout"] = connection_timeout 

258 

259 if self.statement_cache_size is not None: 259 ↛ 260line 259 didn't jump to line 260 because the condition on line 259 was never true

260 connect_args["statement_cache_size"] = self.statement_cache_size 

261 

262 if self.prepared_statement_cache_size is not None: 262 ↛ 263line 262 didn't jump to line 263 because the condition on line 262 was never true

263 connect_args["prepared_statement_cache_size"] = ( 

264 self.prepared_statement_cache_size 

265 ) 

266 

267 server_settings: dict[str, str] = {} 

268 if self.connection_app_name is not None: 268 ↛ 269line 268 didn't jump to line 269 because the condition on line 268 was never true

269 server_settings["application_name"] = self.connection_app_name 

270 if self.search_path is not None: 270 ↛ 271line 270 didn't jump to line 271 because the condition on line 270 was never true

271 server_settings["search_path"] = self.search_path 

272 if server_settings: 272 ↛ 273line 272 didn't jump to line 273 because the condition on line 272 was never true

273 connect_args["server_settings"] = server_settings 

274 

275 if get_current_settings().server.database.sqlalchemy.connect_args.tls.enabled: 275 ↛ 276line 275 didn't jump to line 276 because the condition on line 275 was never true

276 tls_config = ( 

277 get_current_settings().server.database.sqlalchemy.connect_args.tls 

278 ) 

279 

280 pg_ctx = ssl.create_default_context(purpose=ssl.Purpose.SERVER_AUTH) 

281 

282 if tls_config.ca_file: 282 ↛ 287line 282 didn't jump to line 287 because the condition on line 282 was always true

283 pg_ctx = ssl.create_default_context( 

284 purpose=ssl.Purpose.SERVER_AUTH, cafile=tls_config.ca_file 

285 ) 

286 

287 pg_ctx.minimum_version = ssl.TLSVersion.TLSv1_2 

288 

289 if tls_config.cert_file and tls_config.key_file: 

290 pg_ctx.load_cert_chain( 

291 certfile=tls_config.cert_file, keyfile=tls_config.key_file 

292 ) 

293 

294 pg_ctx.check_hostname = tls_config.check_hostname 

295 pg_ctx.verify_mode = ssl.CERT_REQUIRED 

296 connect_args["ssl"] = pg_ctx 

297 

298 # Initialize plugin manager 

299 if get_current_settings().plugins.enabled: 299 ↛ 300line 299 didn't jump to line 300 because the condition on line 299 was never true

300 pm = build_manager(HookSpec) 

301 load_entry_point_plugins( 

302 pm, 

303 allow=get_current_settings().plugins.allow, 

304 deny=get_current_settings().plugins.deny, 

305 logger=logging.getLogger("prefect.plugins"), 

306 ) 

307 

308 # Call set_database_connection_params hook 

309 results = await call_async_hook( 

310 pm, 

311 "set_database_connection_params", 

312 connection_url=self.connection_url, 

313 settings=get_current_settings(), 

314 ) 

315 

316 for _, params, error in results: 

317 if error: 

318 # Log error but don't fail, other plugins might succeed 

319 logging.getLogger("prefect.server.database").warning( 

320 "Plugin failed to set database connection params: %s", error 

321 ) 

322 elif params: 

323 connect_args.update(params) 

324 

325 if connect_args: 325 ↛ 328line 325 didn't jump to line 328 because the condition on line 325 was always true

326 kwargs["connect_args"] = connect_args 

327 

328 if self.sqlalchemy_pool_size is not None: 328 ↛ 331line 328 didn't jump to line 331 because the condition on line 328 was always true

329 kwargs["pool_size"] = self.sqlalchemy_pool_size 

330 

331 if self.sqlalchemy_max_overflow is not None: 331 ↛ 334line 331 didn't jump to line 334 because the condition on line 331 was always true

332 kwargs["max_overflow"] = self.sqlalchemy_max_overflow 

333 

334 engine = create_async_engine( 

335 self.connection_url, 

336 echo=self.echo, 

337 # "pre-ping" connections upon checkout to ensure they have not been 

338 # closed on the server side 

339 pool_pre_ping=True, 

340 # Use connections in LIFO order to help reduce connections 

341 # after spiky load and in general increase the likelihood 

342 # that a given connection pulled from the pool will be 

343 # usable. 

344 pool_use_lifo=True, 

345 **kwargs, 

346 ) 

347 

348 if logfire: 348 ↛ 349line 348 didn't jump to line 349 because the condition on line 348 was never true

349 logfire.instrument_sqlalchemy(engine) # pyright: ignore 

350 

351 if TRACKER.active: 351 ↛ 352line 351 didn't jump to line 352 because the condition on line 351 was never true

352 TRACKER.track_pool(engine.pool) 

353 

354 ENGINES[cache_key] = engine 

355 await self.schedule_engine_disposal(cache_key) 

356 return ENGINES[cache_key] 

357 

358 async def schedule_engine_disposal(self, cache_key: _EngineCacheKey) -> None: 

359 """ 

360 Dispose of an engine once the event loop is closing. 

361 

362 See caveats at `add_event_loop_shutdown_callback`. 

363 

364 We attempted to lazily clean up old engines when new engines are created, but 

365 if the loop the engine is attached to is already closed then the connections 

366 cannot be cleaned up properly and warnings are displayed. 

367 

368 Engine disposal should only be important when running the application 

369 ephemerally. Notably, this is an issue in our tests where many short-lived event 

370 loops and engines are created which can consume all of the available database 

371 connection slots. Users operating at a scale where connection limits are 

372 encountered should be encouraged to use a standalone server. 

373 """ 

374 

375 async def dispose_engine(cache_key: _EngineCacheKey) -> None: 

376 engine = ENGINES.pop(cache_key, None) 

377 if engine: 

378 await engine.dispose() 

379 

380 await add_event_loop_shutdown_callback(partial(dispose_engine, cache_key)) 

381 

382 async def session(self, engine: AsyncEngine) -> AsyncSession: 

383 """ 

384 Retrieves a SQLAlchemy session for an engine. 

385 

386 Args: 

387 engine: a sqlalchemy engine 

388 """ 

389 return AsyncSession(engine, expire_on_commit=False) 

390 

391 @asynccontextmanager 

392 async def begin_transaction( 

393 self, session: AsyncSession, with_for_update: bool = False 

394 ) -> AsyncGenerator[AsyncSessionTransaction, None]: 

395 # `with_for_update` is for SQLite only. For Postgres, lock the row on read 

396 # for update instead. 

397 async with session.begin() as transaction: 

398 yield transaction 

399 

400 async def create_db( 

401 self, connection: AsyncConnection, base_metadata: sa.MetaData 

402 ) -> None: 

403 """Create the database""" 

404 

405 await connection.run_sync(base_metadata.create_all) 

406 

407 async def drop_db( 

408 self, connection: AsyncConnection, base_metadata: sa.MetaData 

409 ) -> None: 

410 """Drop the database""" 

411 

412 await connection.run_sync(base_metadata.drop_all) 

413 

414 def is_inmemory(self) -> bool: 

415 """Returns true if database is run in memory""" 

416 

417 return False 

418 

419 

420class AioSqliteConfiguration(BaseDatabaseConfiguration): 

421 MIN_SQLITE_VERSION = (3, 24, 0) 

422 

423 async def engine(self) -> AsyncEngine: 

424 """Retrieves an async SQLAlchemy engine. 

425 

426 Args: 

427 connection_url (str, optional): The database connection string. 

428 Defaults to self.connection_url 

429 echo (bool, optional): Whether to echo SQL sent 

430 to the database. Defaults to self.echo 

431 timeout (float, optional): The database statement timeout, in seconds. 

432 Defaults to self.timeout 

433 

434 Returns: 

435 AsyncEngine: a SQLAlchemy engine 

436 """ 

437 

438 if sqlite3.sqlite_version_info < self.MIN_SQLITE_VERSION: 

439 required = ".".join(str(v) for v in self.MIN_SQLITE_VERSION) 

440 raise RuntimeError( 

441 f"Prefect requires sqlite >= {required} but we found version " 

442 f"{sqlite3.sqlite_version}" 

443 ) 

444 

445 kwargs: dict[str, Any] = dict() 

446 

447 loop = get_running_loop() 

448 

449 cache_key = (loop, self.connection_url, self.echo, self.timeout) 

450 if cache_key not in ENGINES: 

451 # apply database timeout 

452 # In test mode, use a higher timeout to handle lock contention during 

453 # parallel test execution. This should match the PRAGMA busy_timeout. 

454 if PREFECT_TESTING_UNIT_TEST_MODE.value() is True: 

455 kwargs["connect_args"] = dict(timeout=30.0) # 30s for tests 

456 elif self.timeout is not None: 456 ↛ anywhereline 456 didn't jump anywhere: it always raised an exception.

457 kwargs["connect_args"] = dict(timeout=self.timeout) 

458 

459 # use `named` paramstyle for sqlite instead of `qmark` in very rare 

460 # circumstances, we've seen aiosqlite pass parameters in the wrong 

461 # order; by using named parameters we avoid this issue 

462 # see https://github.com/PrefectHQ/prefect/pull/6702 

463 kwargs["paramstyle"] = "named" 

464 

465 # ensure a long-lasting pool is used with in-memory databases 

466 # because they disappear when the last connection closes 

467 if ":memory:" in self.connection_url: 

468 kwargs.update( 

469 poolclass=sa.pool.AsyncAdaptedQueuePool, 

470 pool_size=1, 

471 max_overflow=0, 

472 pool_recycle=-1, 

473 ) 

474 

475 engine = create_async_engine(self.connection_url, echo=self.echo, **kwargs) 

476 event.listen(engine.sync_engine, "connect", self.setup_sqlite) 

477 event.listen(engine.sync_engine, "begin", self.begin_sqlite_stmt) 

478 

479 if logfire: 479 ↛ 482line 479 didn't jump to line 482 because the condition on line 479 was always true

480 logfire.instrument_sqlalchemy(engine) # pyright: ignore 

481 

482 if TRACKER.active: 

483 TRACKER.track_pool(engine.pool) 

484 

485 ENGINES[cache_key] = engine 

486 await self.schedule_engine_disposal(cache_key) 

487 return ENGINES[cache_key] 

488 

489 async def schedule_engine_disposal(self, cache_key: _EngineCacheKey) -> None: 

490 """ 

491 Dispose of an engine once the event loop is closing. 

492 

493 See caveats at `add_event_loop_shutdown_callback`. 

494 

495 We attempted to lazily clean up old engines when new engines are created, but 

496 if the loop the engine is attached to is already closed then the connections 

497 cannot be cleaned up properly and warnings are displayed. 

498 

499 Engine disposal should only be important when running the application 

500 ephemerally. Notably, this is an issue in our tests where many short-lived event 

501 loops and engines are created which can consume all of the available database 

502 connection slots. Users operating at a scale where connection limits are 

503 encountered should be encouraged to use a standalone server. 

504 """ 

505 

506 async def dispose_engine(cache_key: _EngineCacheKey) -> None: 

507 engine = ENGINES.pop(cache_key, None) 

508 if engine: 

509 await engine.dispose() 

510 

511 await add_event_loop_shutdown_callback(partial(dispose_engine, cache_key)) 

512 

513 def setup_sqlite(self, conn: DBAPIConnection, record: ConnectionPoolEntry) -> None: 

514 """Issue PRAGMA statements to SQLITE on connect. PRAGMAs only last for the 

515 duration of the connection. See https://www.sqlite.org/pragma.html for more info. 

516 """ 

517 # workaround sqlite transaction behavior 

518 if isinstance(conn, aiosqlite.AsyncAdapt_aiosqlite_connection): 

519 self.begin_sqlite_conn(conn) 

520 

521 cursor = conn.cursor() 

522 

523 # write to a write-ahead-log instead and regularly commit the changes 

524 # this allows multiple concurrent readers even during a write transaction 

525 # even with the WAL we can get busy errors if we have transactions that: 

526 # - t1 reads from a database 

527 # - t2 inserts to the database 

528 # - t1 tries to insert to the database 

529 # this can be resolved by using the IMMEDIATE transaction mode in t1 

530 cursor.execute("PRAGMA journal_mode = WAL;") 

531 

532 # enable foreign keys 

533 cursor.execute("PRAGMA foreign_keys = ON;") 

534 

535 # disable legacy alter table behavior as it will cause problems during 

536 # migrations when tables are renamed as references would otherwise be retained 

537 # in some locations 

538 # https://www.sqlite.org/pragma.html#pragma_legacy_alter_table 

539 cursor.execute("PRAGMA legacy_alter_table=OFF") 

540 

541 # when using the WAL, we do need to sync changes on every write. sqlite 

542 # recommends using 'normal' mode which is much faster 

543 cursor.execute("PRAGMA synchronous = NORMAL;") 

544 

545 # a higher cache size (default of 2000) for more aggressive performance 

546 cursor.execute("PRAGMA cache_size = 20000;") 

547 

548 # wait for this amount of time while a table is locked 

549 # before returning and raising an error 

550 # setting the value very high allows for more 'concurrency' 

551 # without running into errors, but may result in slow api calls 

552 if PREFECT_TESTING_UNIT_TEST_MODE.value() is True: 

553 cursor.execute("PRAGMA busy_timeout = 30000;") # 30s 

554 else: 

555 cursor.execute("PRAGMA busy_timeout = 60000;") # 60s 

556 

557 # `PRAGMA temp_store = memory;` moves temporary tables from disk into RAM 

558 # this supposedly speeds up reads, but it seems to actually 

559 # decrease overall performance, see https://github.com/PrefectHQ/prefect/pull/14812 

560 # cursor.execute("PRAGMA temp_store = memory;") 

561 

562 cursor.close() 

563 

564 def begin_sqlite_conn( 

565 self, conn: aiosqlite.AsyncAdapt_aiosqlite_connection 

566 ) -> None: 

567 # disable pysqlite's emitting of the BEGIN statement entirely. 

568 # also stops it from emitting COMMIT before any DDL. 

569 # requires `begin_sqlite_stmt` 

570 # see https://docs.sqlalchemy.org/en/20/dialects/sqlite.html#serializable-isolation-savepoints-transactional-ddl 

571 conn.isolation_level = None 

572 

573 def begin_sqlite_stmt(self, conn: sa.Connection) -> None: 

574 # emit our own BEGIN 

575 # requires `begin_sqlite_conn` 

576 # see https://docs.sqlalchemy.org/en/20/dialects/sqlite.html#serializable-isolation-savepoints-transactional-ddl 

577 mode = SQLITE_BEGIN_MODE.get() 

578 if mode is not None: 

579 conn.exec_driver_sql(f"BEGIN {mode}") 

580 

581 # Note this is intentionally a no-op if there is no BEGIN MODE set 

582 # This allows us to use SQLite's default behavior for reads which do not need 

583 # to be wrapped in a long-running transaction 

584 

585 @asynccontextmanager 

586 async def begin_transaction( 

587 self, session: AsyncSession, with_for_update: bool = False 

588 ) -> AsyncGenerator[AsyncSessionTransaction, None]: 

589 token = SQLITE_BEGIN_MODE.set("IMMEDIATE" if with_for_update else "DEFERRED") 

590 

591 try: 

592 async with session.begin() as transaction: 

593 yield transaction 

594 finally: 

595 SQLITE_BEGIN_MODE.reset(token) 

596 

597 async def session(self, engine: AsyncEngine) -> AsyncSession: 

598 """ 

599 Retrieves a SQLAlchemy session for an engine. 

600 

601 Args: 

602 engine: a sqlalchemy engine 

603 """ 

604 return AsyncSession(engine, expire_on_commit=False) 

605 

606 async def create_db( 

607 self, connection: AsyncConnection, base_metadata: sa.MetaData 

608 ) -> None: 

609 """Create the database""" 

610 

611 await connection.run_sync(base_metadata.create_all) 

612 

613 async def drop_db( 

614 self, connection: AsyncConnection, base_metadata: sa.MetaData 

615 ) -> None: 

616 """Drop the database""" 

617 

618 await connection.run_sync(base_metadata.drop_all) 

619 

620 def is_inmemory(self) -> bool: 

621 """Returns true if database is run in memory""" 

622 

623 return ":memory:" in self.connection_url or "mode=memory" in self.connection_url