session.py 2.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081
  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, event, 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 _set_mysql_session_utc(dbapi_connection: Any, _connection_record: Any) -> None:
  12. """Keep MySQL-generated timestamps aligned with application UTC values."""
  13. cursor = dbapi_connection.cursor()
  14. try:
  15. cursor.execute("SET time_zone = '+00:00'")
  16. finally:
  17. cursor.close()
  18. def get_engine():
  19. """Lazy-init SQLAlchemy engine (singleton)."""
  20. global _engine, _SessionLocal
  21. if _engine is None:
  22. settings = get_infra_settings()
  23. _engine = create_engine(
  24. settings.mysql_url,
  25. pool_size=settings.selected_mysql_pool_size,
  26. max_overflow=settings.mysql_max_overflow,
  27. pool_timeout=settings.mysql_pool_timeout_seconds,
  28. pool_recycle=settings.mysql_pool_recycle_seconds,
  29. pool_pre_ping=True,
  30. echo=settings.mysql_echo,
  31. )
  32. if _engine.dialect.name == "mysql":
  33. event.listen(_engine, "connect", _set_mysql_session_utc)
  34. _SessionLocal = sessionmaker(bind=_engine, autoflush=False, autocommit=False)
  35. return _engine
  36. def dispose_engine() -> None:
  37. """Dispose process-local connections, mainly for graceful shutdown and tests."""
  38. global _engine, _SessionLocal
  39. if _engine is not None:
  40. _engine.dispose()
  41. _engine = None
  42. _SessionLocal = None
  43. def init_db() -> dict[str, list[str]]:
  44. """Create all tables (dev / first-run). Import models before calling."""
  45. import supply_infra.db.models # noqa: F401 — register all models
  46. engine = get_engine()
  47. inspector = inspect(engine)
  48. before = set(inspector.get_table_names())
  49. Base.metadata.create_all(bind=engine)
  50. after = set(inspect(engine).get_table_names())
  51. created = sorted(after - before)
  52. return {"created": created}
  53. @contextmanager
  54. def get_session() -> Generator[Session, None, None]:
  55. """Provide a transactional database session."""
  56. get_engine()
  57. assert _SessionLocal is not None
  58. session = _SessionLocal()
  59. try:
  60. yield session
  61. session.commit()
  62. except Exception:
  63. session.rollback()
  64. raise
  65. finally:
  66. session.close()