summaryrefslogtreecommitdiff
path: root/src/srht_contrib/db.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/srht_contrib/db.py')
-rw-r--r--src/srht_contrib/db.py23
1 files changed, 20 insertions, 3 deletions
diff --git a/src/srht_contrib/db.py b/src/srht_contrib/db.py
index d96de57..7750df6 100644
--- a/src/srht_contrib/db.py
+++ b/src/srht_contrib/db.py
@@ -3,7 +3,7 @@ from __future__ import annotations
from collections.abc import Generator
from fastapi import HTTPException, Request, status
-from sqlalchemy import Engine, create_engine, text
+from sqlalchemy import Engine, create_engine, event, text
from sqlalchemy.pool import StaticPool
from sqlalchemy.orm import Session, declarative_base, sessionmaker
@@ -13,11 +13,28 @@ Base = declarative_base()
def make_engine(settings: Settings) -> Engine:
- connect_args = {"check_same_thread": False} if settings.database_url.startswith("sqlite") else {}
+ connect_args = {}
+ if settings.database_url.startswith("sqlite"):
+ connect_args = {
+ "check_same_thread": False,
+ "timeout": settings.sqlite_busy_timeout_seconds,
+ }
engine_kwargs = {"future": True, "connect_args": connect_args}
if settings.database_url in {"sqlite://", "sqlite:///:memory:"}:
engine_kwargs["poolclass"] = StaticPool
- return create_engine(settings.database_url, **engine_kwargs)
+ engine = create_engine(settings.database_url, **engine_kwargs)
+
+ if settings.database_url.startswith("sqlite"):
+ @event.listens_for(engine, "connect")
+ def _configure_sqlite(dbapi_connection, connection_record) -> None: # type: ignore[unused-ignore]
+ cursor = dbapi_connection.cursor()
+ cursor.execute(f"PRAGMA busy_timeout = {int(settings.sqlite_busy_timeout_seconds * 1000)}")
+ if settings.database_url not in {"sqlite://", "sqlite:///:memory:"}:
+ cursor.execute("PRAGMA journal_mode = WAL")
+ cursor.execute("PRAGMA synchronous = NORMAL")
+ cursor.close()
+
+ return engine
def make_session_factory(settings: Settings) -> sessionmaker[Session]: