session.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120
  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. # Process-local pool override. Used by parallel step workers so we can grow
  12. # this process's pool without mutating MYSQL_POOL_SIZE_CONTROL (which would
  13. # re-validate and explode the fleet-wide connection budget).
  14. _pool_size_override: int | None = None
  15. def _set_mysql_session_china_time(
  16. dbapi_connection: Any,
  17. _connection_record: Any,
  18. ) -> None:
  19. """Store MySQL-generated timestamps as China Standard Time."""
  20. cursor = dbapi_connection.cursor()
  21. try:
  22. cursor.execute("SET time_zone = '+08:00'")
  23. finally:
  24. cursor.close()
  25. def _resolve_pool_size() -> int:
  26. if _pool_size_override is not None:
  27. return max(1, int(_pool_size_override))
  28. return get_infra_settings().selected_mysql_pool_size
  29. def get_engine():
  30. """Lazy-init SQLAlchemy engine (singleton)."""
  31. global _engine, _SessionLocal
  32. if _engine is None:
  33. settings = get_infra_settings()
  34. _engine = create_engine(
  35. settings.mysql_url,
  36. pool_size=_resolve_pool_size(),
  37. max_overflow=settings.mysql_max_overflow,
  38. pool_timeout=settings.mysql_pool_timeout_seconds,
  39. pool_recycle=settings.mysql_pool_recycle_seconds,
  40. pool_pre_ping=True,
  41. echo=settings.mysql_echo,
  42. # pymysql 默认 read/write timeout 为 None(永不超时):一旦锁等待或网络抖动,
  43. # 查询会在 socket.recv() 上永久阻塞,且这类同步阻塞无法被 asyncio 超时取消。
  44. # 这里显式加 socket 级超时,确保任何一次查询最多阻塞有限时间就会抛异常。
  45. connect_args={
  46. "connect_timeout": settings.mysql_connect_timeout_seconds,
  47. "read_timeout": settings.mysql_read_timeout_seconds,
  48. "write_timeout": settings.mysql_write_timeout_seconds,
  49. },
  50. )
  51. if _engine.dialect.name == "mysql":
  52. event.listen(_engine, "connect", _set_mysql_session_china_time)
  53. _SessionLocal = sessionmaker(bind=_engine, autoflush=False, autocommit=False)
  54. return _engine
  55. def dispose_engine() -> None:
  56. """Dispose process-local connections, mainly for graceful shutdown and tests."""
  57. global _engine, _SessionLocal
  58. if _engine is not None:
  59. _engine.dispose()
  60. _engine = None
  61. _SessionLocal = None
  62. def ensure_mysql_pool_capacity(min_connections: int) -> None:
  63. """Grow the process-local pool before parallel DB access in the same process.
  64. Must not mutate MYSQL_POOL_SIZE_CONTROL: that setting feeds the fleet-wide
  65. connection budget validator. Raising it for a single step (e.g. to 5) makes
  66. required connections jump to 60 against a 40 budget and crashes the step.
  67. """
  68. global _pool_size_override
  69. required = max(1, int(min_connections))
  70. if _resolve_pool_size() >= required:
  71. return
  72. _pool_size_override = required
  73. dispose_engine()
  74. def init_db() -> dict[str, list[str]]:
  75. """Create all tables (dev / first-run). Import models before calling."""
  76. import supply_infra.db.models # noqa: F401 — register all models
  77. import find_agent_v2.models # noqa: F401 — isolated find_agent_v2 tables
  78. engine = get_engine()
  79. inspector = inspect(engine)
  80. before = set(inspector.get_table_names())
  81. Base.metadata.create_all(bind=engine)
  82. after = set(inspect(engine).get_table_names())
  83. created = sorted(after - before)
  84. return {"created": created}
  85. @contextmanager
  86. def get_session() -> Generator[Session, None, None]:
  87. """Provide a transactional database session."""
  88. get_engine()
  89. assert _SessionLocal is not None
  90. session = _SessionLocal()
  91. try:
  92. yield session
  93. session.commit()
  94. except Exception:
  95. session.rollback()
  96. raise
  97. finally:
  98. session.close()