session.py 1.6 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758
  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, inspect
  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() -> dict[str, list[str]]:
  25. """Create all tables (dev / first-run). Import models before calling."""
  26. import supply_infra.db.models # noqa: F401 — register all models
  27. engine = get_engine()
  28. inspector = inspect(engine)
  29. before = set(inspector.get_table_names())
  30. Base.metadata.create_all(bind=engine)
  31. after = set(inspect(engine).get_table_names())
  32. created = sorted(after - before)
  33. return {"created": created}
  34. @contextmanager
  35. def get_session() -> Generator[Session, None, None]:
  36. """Provide a transactional database session."""
  37. get_engine()
  38. assert _SessionLocal is not None
  39. session = _SessionLocal()
  40. try:
  41. yield session
  42. session.commit()
  43. except Exception:
  44. session.rollback()
  45. raise
  46. finally:
  47. session.close()