from collections.abc import Generator from pathlib import Path from sqlalchemy import Engine, create_engine, event from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker from app.config import get_settings class Base(DeclarativeBase): pass def create_sqlite_engine(database_url: str): if database_url.startswith("sqlite:///") and database_url != "sqlite:///:memory:": db_path = Path(database_url.removeprefix("sqlite:///")) if db_path.parent != Path("."): db_path.parent.mkdir(parents=True, exist_ok=True) connect_args = {"check_same_thread": False, "timeout": 10} engine = create_engine(database_url, connect_args=connect_args) engine.dialect.connect_args = connect_args @event.listens_for(engine, "connect") def set_sqlite_pragma(dbapi_connection, _connection_record): cursor = dbapi_connection.cursor() cursor.execute("PRAGMA journal_mode=WAL") cursor.close() return engine engine = create_sqlite_engine(get_settings().database_url) SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False) def check_database_integrity(target_engine: Engine = engine) -> str: raw_connection = target_engine.raw_connection() try: cursor = raw_connection.driver_connection.cursor() try: cursor.execute("PRAGMA integrity_check") result = cursor.fetchone() finally: cursor.close() finally: raw_connection.close() status = result[0] if result else "missing_result" if status != "ok": raise RuntimeError(f"SQLite integrity check failed: {status}") return status def checkpoint_sqlite_wal(target_engine: Engine = engine) -> bool: if str(target_engine.url) == "sqlite:///:memory:": return False if target_engine.url.get_backend_name() != "sqlite": return False with target_engine.begin() as connection: connection.exec_driver_sql("PRAGMA wal_checkpoint(TRUNCATE)") return True def init_db() -> None: import app.models # noqa: F401 Base.metadata.create_all(engine) check_database_integrity(engine) checkpoint_sqlite_wal(engine) def get_db_session() -> Generator[Session]: with SessionLocal() as session: yield session