connection.py 3.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112
  1. """数据库连接管理
  2. 提供统一的数据库连接接口,支持环境变量配置
  3. 参考 content_finder/db/connection.py
  4. """
  5. import os
  6. import logging
  7. from typing import Optional
  8. try:
  9. import pymysql
  10. import pymysql.cursors
  11. except ImportError:
  12. pymysql = None
  13. logger = logging.getLogger(__name__)
  14. def get_connection():
  15. """获取数据库连接
  16. 从环境变量读取配置:
  17. - DB_HOST: 数据库主机地址
  18. - DB_PORT: 数据库端口(默认3306)
  19. - DB_USER: 数据库用户名
  20. - DB_PASSWORD: 数据库密码
  21. - DB_NAME: 数据库名称
  22. Returns:
  23. pymysql.Connection: 数据库连接对象
  24. Raises:
  25. ImportError: pymysql 未安装
  26. ValueError: 数据库配置缺失
  27. Exception: 连接失败
  28. """
  29. if pymysql is None:
  30. raise ImportError(
  31. "pymysql 未安装,请运行: pip install pymysql\n"
  32. "或在 requirements.txt 中添加 pymysql>=1.0.0"
  33. )
  34. # 读取环境变量
  35. host = os.getenv("DB_HOST", "").strip()
  36. port = int(os.getenv("DB_PORT", "3306"))
  37. user = os.getenv("DB_USER", "").strip()
  38. password = os.getenv("DB_PASSWORD", "")
  39. database = os.getenv("DB_NAME", "").strip()
  40. # 验证必需配置
  41. if not all([host, user, database]):
  42. raise ValueError(
  43. "数据库配置缺失!请在 .env 文件或环境变量中设置:\n"
  44. " DB_HOST=数据库主机地址\n"
  45. " DB_USER=数据库用户名\n"
  46. " DB_PASSWORD=数据库密码\n"
  47. " DB_NAME=数据库名称\n"
  48. " DB_PORT=3306 # 可选,默认3306"
  49. )
  50. connect_timeout = int(os.getenv("DB_CONNECT_TIMEOUT", "10"))
  51. read_timeout = int(os.getenv("DB_READ_TIMEOUT", "60"))
  52. write_timeout = int(os.getenv("DB_WRITE_TIMEOUT", "60"))
  53. try:
  54. conn = pymysql.connect(
  55. host=host,
  56. port=port,
  57. user=user,
  58. password=password,
  59. database=database,
  60. charset="utf8mb4",
  61. cursorclass=pymysql.cursors.DictCursor, # 返回字典格式
  62. autocommit=True, # 自动提交(简化事务管理)
  63. connect_timeout=connect_timeout,
  64. read_timeout=read_timeout,
  65. write_timeout=write_timeout,
  66. )
  67. logger.debug(f"数据库连接成功: {user}@{host}:{port}/{database}")
  68. return conn
  69. except Exception as e:
  70. logger.error(f"数据库连接失败: {e}")
  71. raise
  72. def test_connection() -> bool:
  73. """测试数据库连接
  74. Returns:
  75. bool: 连接成功返回 True,失败返回 False
  76. """
  77. try:
  78. conn = get_connection()
  79. with conn.cursor() as cursor:
  80. cursor.execute("SELECT 1")
  81. result = cursor.fetchone()
  82. conn.close()
  83. logger.info("数据库连接测试成功")
  84. return True
  85. except Exception as e:
  86. logger.error(f"数据库连接测试失败: {e}")
  87. return False
  88. if __name__ == "__main__":
  89. # 测试数据库连接
  90. logging.basicConfig(level=logging.INFO)
  91. if test_connection():
  92. print("✅ 数据库连接正常")
  93. else:
  94. print("❌ 数据库连接失败,请检查配置")