from collections.abc import Iterator from sqlalchemy import create_engine, event, inspect, text from sqlalchemy.engine import Engine from sqlalchemy.orm import Session, sessionmaker from app.db.models import Base def create_session_factory(database_url: str) -> tuple[Engine, sessionmaker[Session]]: engine = create_engine(database_url, connect_args={"check_same_thread": False}) @event.listens_for(engine, "connect") def enable_sqlite_foreign_keys(dbapi_connection, _connection_record) -> None: cursor = dbapi_connection.cursor() cursor.execute("PRAGMA foreign_keys=ON") cursor.close() return engine, sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) def init_db(engine) -> None: Base.metadata.create_all(bind=engine) _add_missing_project_columns(engine) def _add_missing_project_columns(engine: Engine) -> None: inspector = inspect(engine) if "projects" not in inspector.get_table_names(): return columns = {column["name"] for column in inspector.get_columns("projects")} if "drug_name" in columns: return with engine.begin() as connection: connection.execute(text("ALTER TABLE projects ADD COLUMN drug_name TEXT")) def get_session(session_factory) -> Iterator[Session]: with session_factory() as session: yield session