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