aboutsummaryrefslogtreecommitdiffstats
path: root/packages/meshbay-hub/src/meshbay_hub/db/engine.py
diff options
context:
space:
mode:
Diffstat (limited to 'packages/meshbay-hub/src/meshbay_hub/db/engine.py')
-rw-r--r--packages/meshbay-hub/src/meshbay_hub/db/engine.py78
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