db.py 1.8 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364
  1. # -*- coding: utf-8 -*-
  2. """统一 PostgreSQL 连接池。
  3. 合并原 common/Connection.py 与 interface/Connection.py(两者完全一致)。
  4. 所有 PG 连接经此模块,配置走 infra/config.py。
  5. 向后兼容:
  6. - common/Connection.getConnection() 继续可用
  7. - interface/Connection.getConnection() 继续可用
  8. 二者均 re-export 自本模块。
  9. """
  10. from __future__ import absolute_import
  11. import logging
  12. import psycopg2
  13. from DBUtils.PooledDB import PooledDB
  14. from . import config
  15. __all__ = ["get_connection", "get_db_connection", "get_pool"]
  16. logger = logging.getLogger(__name__)
  17. _pools = {}
  18. def get_pool(dbname="BiddingKG"):
  19. """获取指定库的连接池,按 dbname 缓存。"""
  20. if dbname not in _pools:
  21. cfg = config.pg_db_config(dbname)
  22. _pools[dbname] = PooledDB(
  23. psycopg2,
  24. int(cfg.get("pool_size", 10)),
  25. host=cfg.get("host", "127.0.0.1"),
  26. port=str(cfg.get("port", 5432)),
  27. user=cfg.get("user", "postgres"),
  28. password=cfg.get("password", ""),
  29. dbname=cfg.get("dbname", dbname),
  30. )
  31. return _pools[dbname]
  32. def get_connection(dbname="BiddingKG"):
  33. """从连接池取一个连接。默认取 BiddingKG 库。
  34. 向后兼容:等价于原 common/Connection.getConnection()。
  35. """
  36. return get_pool(dbname).connection()
  37. def get_db_connection(dbname):
  38. """显式按库名取连接。用于训练/数据脚本访问非默认库。"""
  39. return get_connection(dbname)
  40. def close_all():
  41. """关闭所有连接池,用于测试与进程退出。"""
  42. for dbname, pool in _pools.items():
  43. try:
  44. pool.close()
  45. except Exception as e: # pragma: no cover
  46. logger.warning("close pool %s failed: %s", dbname, e)
  47. _pools.clear()