app/db/__init__.py (view raw)
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 |
from __future__ import annotations
from collections.abc import AsyncGenerator
from pathlib import Path
from sqlalchemy import event, text
from sqlalchemy.ext.asyncio import (
AsyncEngine,
AsyncSession,
async_sessionmaker,
create_async_engine,
)
from app.config import Settings, get_settings
_engine: AsyncEngine | None = None
_session_factory: async_sessionmaker[AsyncSession] | None = None
def _ensure_sqlite_parent(database_url: str) -> None:
if "sqlite" not in database_url:
return
# sqlite+aiosqlite:///data/rastro.db or ////absolute
raw = database_url.split("///", 1)[-1]
path = Path(raw)
if path.parent and str(path.parent) not in {".", ""}:
path.parent.mkdir(parents=True, exist_ok=True)
def get_engine(settings: Settings | None = None) -> AsyncEngine:
global _engine, _session_factory
if _engine is not None:
return _engine
settings = settings or get_settings()
_ensure_sqlite_parent(settings.database_url)
_engine = create_async_engine(
settings.database_url,
echo=False,
connect_args={"check_same_thread": False},
)
@event.listens_for(_engine.sync_engine, "connect")
def _set_sqlite_pragma(dbapi_connection, connection_record) -> None: # noqa: ARG001
cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA journal_mode=WAL")
cursor.execute("PRAGMA foreign_keys=ON")
cursor.close()
_session_factory = async_sessionmaker(_engine, expire_on_commit=False)
return _engine
def get_session_factory() -> async_sessionmaker[AsyncSession]:
if _session_factory is None:
get_engine()
assert _session_factory is not None
return _session_factory
async def get_session() -> AsyncGenerator[AsyncSession]:
factory = get_session_factory()
async with factory() as session:
yield session
async def init_db(settings: Settings | None = None) -> None:
"""Create tables if needed (Alembic preferred; used as fallback/tests)."""
from app.db.models import Base
engine = get_engine(settings)
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
await conn.execute(text("PRAGMA journal_mode=WAL"))
def reset_engine() -> None:
"""Reset global engine (tests)."""
global _engine, _session_factory
_engine = None
_session_factory = None
|