from contextlib import contextmanager from pathlib import Path from typing import Generator from sqlalchemy import create_engine, event, text from sqlalchemy.orm import Session, sessionmaker from gimme_job.models.db import Base def get_engine(db_path: Path | None = None): from gimme_job.utils.paths import db_path as default_db_path path = db_path or default_db_path() path.parent.mkdir(parents=True, exist_ok=True) engine = create_engine(f"sqlite:///{path}", echo=False) # Enable WAL mode and foreign keys @event.listens_for(engine, "connect") def set_sqlite_pragma(dbapi_conn, _connection_record): cursor = dbapi_conn.cursor() cursor.execute("PRAGMA journal_mode=WAL") cursor.execute("PRAGMA foreign_keys=ON") cursor.close() return engine def get_session_factory(engine=None): if engine is None: engine = get_engine() return sessionmaker(bind=engine, autoflush=False, autocommit=False) def init_db(engine=None) -> None: if engine is None: engine = get_engine() Base.metadata.create_all(engine) @contextmanager def db_session(engine=None) -> Generator[Session, None, None]: factory = get_session_factory(engine) session: Session = factory() try: yield session session.commit() except Exception: session.rollback() raise finally: session.close()