| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152 |
- from __future__ import annotations
- from collections.abc import Generator
- from contextlib import contextmanager
- from typing import Any
- from sqlalchemy import create_engine
- 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 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.mysql_pool_size,
- pool_pre_ping=True,
- echo=settings.mysql_echo,
- )
- _SessionLocal = sessionmaker(bind=_engine, autoflush=False, autocommit=False)
- return _engine
- def init_db() -> None:
- """Create all tables (dev / first-run). Import models before calling."""
- import supply_infra.db.models # noqa: F401 — register all models
- Base.metadata.create_all(bind=get_engine())
- @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()
|