session.py 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110
  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. # pymysql 默认 read/write timeout 为 None(永不超时):一旦锁等待或网络抖动,
  35. # 查询会在 socket.recv() 上永久阻塞,且这类同步阻塞无法被 asyncio 超时取消。
  36. # 这里显式加 socket 级超时,确保任何一次查询最多阻塞有限时间就会抛异常。
  37. connect_args={
  38. "connect_timeout": settings.mysql_connect_timeout_seconds,
  39. "read_timeout": settings.mysql_read_timeout_seconds,
  40. "write_timeout": settings.mysql_write_timeout_seconds,
  41. },
  42. )
  43. if _engine.dialect.name == "mysql":
  44. event.listen(_engine, "connect", _set_mysql_session_china_time)
  45. _SessionLocal = sessionmaker(bind=_engine, autoflush=False, autocommit=False)
  46. return _engine
  47. def dispose_engine() -> None:
  48. """Dispose process-local connections, mainly for graceful shutdown and tests."""
  49. global _engine, _SessionLocal
  50. if _engine is not None:
  51. _engine.dispose()
  52. _engine = None
  53. _SessionLocal = None
  54. def ensure_mysql_pool_capacity(min_connections: int) -> None:
  55. """Grow the process-local pool before parallel DB access in the same process."""
  56. required = max(1, int(min_connections))
  57. settings = get_infra_settings()
  58. if settings.selected_mysql_pool_size >= required:
  59. return
  60. import os
  61. role = settings.process_role
  62. if role in {"scheduler", "worker", "reconciler", "step"}:
  63. os.environ["MYSQL_POOL_SIZE_CONTROL"] = str(required)
  64. else:
  65. os.environ["MYSQL_POOL_SIZE"] = str(required)
  66. get_infra_settings.cache_clear()
  67. dispose_engine()
  68. def init_db() -> dict[str, list[str]]:
  69. """Create all tables (dev / first-run). Import models before calling."""
  70. import supply_infra.db.models # noqa: F401 — register all models
  71. engine = get_engine()
  72. inspector = inspect(engine)
  73. before = set(inspector.get_table_names())
  74. Base.metadata.create_all(bind=engine)
  75. after = set(inspect(engine).get_table_names())
  76. created = sorted(after - before)
  77. return {"created": created}
  78. @contextmanager
  79. def get_session() -> Generator[Session, None, None]:
  80. """Provide a transactional database session."""
  81. get_engine()
  82. assert _SessionLocal is not None
  83. session = _SessionLocal()
  84. try:
  85. yield session
  86. session.commit()
  87. except Exception:
  88. session.rollback()
  89. raise
  90. finally:
  91. session.close()