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()