| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384 |
- from __future__ import annotations
- from collections.abc import Generator
- from contextlib import contextmanager
- from typing import Any
- from sqlalchemy import create_engine, event, inspect
- from sqlalchemy.orm import Session, sessionmaker
- from supply_infra.config import get_infra_settings
- from supply_infra.db.base import Base
- _engine: Any = None
- _SessionLocal: sessionmaker[Session] | None = None
- def _set_mysql_session_china_time(
- dbapi_connection: Any,
- _connection_record: Any,
- ) -> None:
- """Store MySQL-generated timestamps as China Standard Time."""
- cursor = dbapi_connection.cursor()
- try:
- cursor.execute("SET time_zone = '+08:00'")
- finally:
- cursor.close()
- def get_engine():
- """Lazy-init SQLAlchemy engine (singleton)."""
- global _engine, _SessionLocal
- if _engine is None:
- settings = get_infra_settings()
- _engine = create_engine(
- settings.mysql_url,
- pool_size=settings.selected_mysql_pool_size,
- max_overflow=settings.mysql_max_overflow,
- pool_timeout=settings.mysql_pool_timeout_seconds,
- pool_recycle=settings.mysql_pool_recycle_seconds,
- pool_pre_ping=True,
- echo=settings.mysql_echo,
- )
- if _engine.dialect.name == "mysql":
- event.listen(_engine, "connect", _set_mysql_session_china_time)
- _SessionLocal = sessionmaker(bind=_engine, autoflush=False, autocommit=False)
- return _engine
- def dispose_engine() -> None:
- """Dispose process-local connections, mainly for graceful shutdown and tests."""
- global _engine, _SessionLocal
- if _engine is not None:
- _engine.dispose()
- _engine = None
- _SessionLocal = None
- def init_db() -> dict[str, list[str]]:
- """Create all tables (dev / first-run). Import models before calling."""
- import supply_infra.db.models # noqa: F401 — register all models
- engine = get_engine()
- inspector = inspect(engine)
- before = set(inspector.get_table_names())
- Base.metadata.create_all(bind=engine)
- after = set(inspect(engine).get_table_names())
- created = sorted(after - before)
- return {"created": created}
- @contextmanager
- def get_session() -> Generator[Session, None, None]:
- """Provide a transactional database session."""
- get_engine()
- assert _SessionLocal is not None
- session = _SessionLocal()
- try:
- yield session
- session.commit()
- except Exception:
- session.rollback()
- raise
- finally:
- session.close()
|