session.py 2.5 KB

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