session.py 1.4 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152
  1. from __future__ import annotations
  2. from collections.abc import Generator
  3. from contextlib import contextmanager
  4. from typing import Any
  5. from sqlalchemy import create_engine
  6. from sqlalchemy.orm import Session, sessionmaker
  7. from supply_infra.config import get_infra_settings
  8. from supply_infra.db.base import Base
  9. _engine: Any = None
  10. _SessionLocal: sessionmaker[Session] | None = None
  11. def get_engine():
  12. """Lazy-init SQLAlchemy engine (singleton)."""
  13. global _engine, _SessionLocal
  14. if _engine is None:
  15. settings = get_infra_settings()
  16. _engine = create_engine(
  17. settings.mysql_url,
  18. pool_size=settings.mysql_pool_size,
  19. pool_pre_ping=True,
  20. echo=settings.mysql_echo,
  21. )
  22. _SessionLocal = sessionmaker(bind=_engine, autoflush=False, autocommit=False)
  23. return _engine
  24. def init_db() -> None:
  25. """Create all tables (dev / first-run). Import models before calling."""
  26. import supply_infra.db.models # noqa: F401 — register all models
  27. Base.metadata.create_all(bind=get_engine())
  28. @contextmanager
  29. def get_session() -> Generator[Session, None, None]:
  30. """Provide a transactional database session."""
  31. get_engine()
  32. assert _SessionLocal is not None
  33. session = _SessionLocal()
  34. try:
  35. yield session
  36. session.commit()
  37. except Exception:
  38. session.rollback()
  39. raise
  40. finally:
  41. session.close()