Skip to content

Commit ea2cf8f

Browse files
authored
Db engine cache refactor (#328)
Use functools to cache the database engine instead of private module variables.
1 parent 9b226b1 commit ea2cf8f

1 file changed

Lines changed: 21 additions & 18 deletions

File tree

src/database/setup.py

Lines changed: 21 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,11 @@
1+
import functools
2+
3+
from loguru import logger
14
from sqlalchemy.engine import URL
25
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
36

47
from config import DatabaseConfiguration, get_config
58

6-
_user_engine = None
7-
_expdb_engine = None
8-
99

1010
def _create_engine(db_config: DatabaseConfiguration) -> AsyncEngine:
1111
db_url = URL.create(
@@ -16,33 +16,36 @@ def _create_engine(db_config: DatabaseConfiguration) -> AsyncEngine:
1616
port=db_config.port,
1717
database=db_config.database,
1818
)
19+
20+
logger.info("Creating database engine for {db_url}", db_url=db_url)
1921
return create_async_engine(
2022
db_url,
2123
echo=db_config.echo,
2224
pool_recycle=3600,
2325
)
2426

2527

28+
@functools.cache
2629
def user_database() -> AsyncEngine:
27-
global _user_engine # noqa: PLW0603
28-
if _user_engine is None:
29-
_user_engine = _create_engine(get_config().openml_database)
30-
return _user_engine
30+
return _create_engine(get_config().openml_database)
3131

3232

33+
@functools.cache
3334
def expdb_database() -> AsyncEngine:
34-
global _expdb_engine # noqa: PLW0603
35-
if _expdb_engine is None:
36-
_expdb_engine = _create_engine(get_config().expdb_database)
37-
return _expdb_engine
35+
return _create_engine(get_config().expdb_database)
3836

3937

4038
async def close_databases() -> None:
4139
"""Close all database connections."""
42-
global _user_engine, _expdb_engine # noqa: PLW0603
43-
if _user_engine is not None:
44-
await _user_engine.dispose()
45-
_user_engine = None
46-
if _expdb_engine is not None:
47-
await _expdb_engine.dispose()
48-
_expdb_engine = None
40+
for db in (user_database, expdb_database):
41+
if db.cache_info().currsize == 1:
42+
engine = db()
43+
logger.info("Disposing of engine connected to {db_url}", db_url=engine.url)
44+
try:
45+
await engine.dispose()
46+
except Exception: # noqa: BLE001
47+
logger.exception(
48+
"Issue disposing of database engine for {db_url}",
49+
db_url=engine.url,
50+
)
51+
db.cache_clear()

0 commit comments

Comments
 (0)