diff options
Diffstat (limited to 'packages/meshbay-hub/src/meshbay_hub/db/engine.py')
| -rw-r--r-- | packages/meshbay-hub/src/meshbay_hub/db/engine.py | 78 |
1 files changed, 78 insertions, 0 deletions
diff --git a/packages/meshbay-hub/src/meshbay_hub/db/engine.py b/packages/meshbay-hub/src/meshbay_hub/db/engine.py new file mode 100644 index 0000000..6ed5238 --- /dev/null +++ b/packages/meshbay-hub/src/meshbay_hub/db/engine.py @@ -0,0 +1,78 @@ +""" +MeshBay Hub — async SQLAlchemy engine and session factory. + +DATABASE_URL env var controls which DB is used: + Production: postgresql+asyncpg://user:pass@localhost/meshbay_hub + Tests: sqlite+aiosqlite:///:memory: (default if not set) +""" + +import os +from collections.abc import AsyncGenerator + +from sqlalchemy.ext.asyncio import ( + AsyncEngine, + AsyncSession, + async_sessionmaker, + create_async_engine, +) + +from meshbay_hub.db.models import Base + +_DEFAULT_URL = "sqlite+aiosqlite:///:memory:" + +def _database_url() -> str: + return os.environ.get("MESHBAY_DATABASE_URL", _DEFAULT_URL) + +# Module-level engine and session factory (initialised in lifespan) +_engine: AsyncEngine | None = None +_session_factory: async_sessionmaker[AsyncSession] | None = None + + +def get_engine() -> AsyncEngine: + if _engine is None: + raise RuntimeError("DB engine not initialised — call init_db() first") + return _engine + + +def get_session_factory() -> async_sessionmaker[AsyncSession]: + if _session_factory is None: + raise RuntimeError("DB not initialised — call init_db() first") + return _session_factory + + +async def init_db(url: str | None = None) -> AsyncEngine: + """Create engine, session factory, and all tables (idempotent).""" + global _engine, _session_factory + + db_url = url or _database_url() + connect_args = {} + if db_url.startswith("sqlite"): + connect_args["check_same_thread"] = False + + _engine = create_async_engine( + db_url, + echo=False, + connect_args=connect_args, + ) + _session_factory = async_sessionmaker( + _engine, expire_on_commit=False, class_=AsyncSession) + + async with _engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + + return _engine + + +async def close_db() -> None: + global _engine, _session_factory + if _engine: + await _engine.dispose() + _engine = None + _session_factory = None + + +async def get_db() -> AsyncGenerator[AsyncSession, None]: + """FastAPI dependency — yields an async DB session.""" + factory = get_session_factory() + async with factory() as session: + yield session |