| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110 |
- 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,
- # pymysql 默认 read/write timeout 为 None(永不超时):一旦锁等待或网络抖动,
- # 查询会在 socket.recv() 上永久阻塞,且这类同步阻塞无法被 asyncio 超时取消。
- # 这里显式加 socket 级超时,确保任何一次查询最多阻塞有限时间就会抛异常。
- connect_args={
- "connect_timeout": settings.mysql_connect_timeout_seconds,
- "read_timeout": settings.mysql_read_timeout_seconds,
- "write_timeout": settings.mysql_write_timeout_seconds,
- },
- )
- 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 ensure_mysql_pool_capacity(min_connections: int) -> None:
- """Grow the process-local pool before parallel DB access in the same process."""
- required = max(1, int(min_connections))
- settings = get_infra_settings()
- if settings.selected_mysql_pool_size >= required:
- return
- import os
- role = settings.process_role
- if role in {"scheduler", "worker", "reconciler", "step"}:
- os.environ["MYSQL_POOL_SIZE_CONTROL"] = str(required)
- else:
- os.environ["MYSQL_POOL_SIZE"] = str(required)
- get_infra_settings.cache_clear()
- dispose_engine()
- 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()
|