""" 数据库连接管理模块 统一管理所有数据库连接池: - asyncpg Pool: 用于一般的数据库操作 - psycopg AsyncConnectionPool: 用于 LangGraph Checkpointer """ from typing import Optional import asyncio import asyncpg import psycopg from psycopg_pool import AsyncConnectionPool from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver from core.config import settings from core.graph_metadata import ensure_graph_metadata, reset_graph_metadata from logger.logging import get_logger logger = get_logger(__name__) # 全局数据库连接池 _asyncpg_pool: Optional[asyncpg.Pool] = None _psycopg_pool: Optional[AsyncConnectionPool] = None _checkpointer: Optional[AsyncPostgresSaver] = None async def _create_asyncpg_pool() -> asyncpg.Pool: """创建 asyncpg 连接池,带指数退避重试。""" max_retries = 3 retry_delay = 2 for attempt in range(max_retries): pool: Optional[asyncpg.Pool] = None try: pool = await asyncpg.create_pool( host=settings.db_host, port=settings.db_port, database=settings.db_name, user=settings.db_user, password=settings.db_password, min_size=settings.db_pool_min_size, max_size=settings.db_pool_max_size, command_timeout=settings.db_command_timeout, timeout=30, max_inactive_connection_lifetime=300, # 闲置 5 分钟后自动丢弃,防止连接腐烂 server_settings={ 'application_name': 'huoyan-enterprise', 'jit': 'off', }, ) async with pool.acquire() as conn: await conn.execute("SELECT 1") await ensure_graph_metadata(conn) logger.info("asyncpg 数据库连接池初始化成功") return pool except Exception as e: logger.error(f"asyncpg 连接池初始化失败 (尝试 {attempt + 1}/{max_retries}): {e}") if pool is not None: try: await pool.close() except Exception: pass if attempt < max_retries - 1: logger.info(f"将在 {retry_delay} 秒后重试...") await asyncio.sleep(retry_delay) retry_delay *= 2 else: logger.error("数据库连接池初始化失败,已达到最大重试次数") raise async def _discard_asyncpg_pool() -> None: """关闭并清除 asyncpg 连接池全局变量,下次调用 get_db_pool 时会重建。""" global _asyncpg_pool if _asyncpg_pool is not None: try: await _asyncpg_pool.close() except Exception: pass _asyncpg_pool = None reset_graph_metadata() logger.info("asyncpg 连接池已重置") async def get_db_pool() -> asyncpg.Pool: """ 获取或创建 asyncpg 数据库连接池。 若检测到连接池已关闭(如数据库重启后连接全部断开), 会自动丢弃旧池并重建,无需重启服务。 """ global _asyncpg_pool # 检查现有 pool 是否已被关闭 if _asyncpg_pool is not None and _asyncpg_pool._closed: logger.warning("检测到 asyncpg 连接池已关闭,将重新初始化") _asyncpg_pool = None reset_graph_metadata() if _asyncpg_pool is None: logger.info( f"初始化 asyncpg 数据库连接池: " f"{settings.db_user}@{settings.db_host}:{settings.db_port}/{settings.db_name}" ) _asyncpg_pool = await _create_asyncpg_pool() return _asyncpg_pool async def _create_psycopg_checkpointer() -> tuple[AsyncConnectionPool, AsyncPostgresSaver]: """创建 psycopg 连接池与 LangGraph Checkpointer,带指数退避重试。""" max_retries = 3 retry_delay = 2 for attempt in range(max_retries): pool: Optional[AsyncConnectionPool] = None try: pool = AsyncConnectionPool( conninfo=settings.db_uri_psycopg, max_size=settings.checkpointer_pool_max_size, open=False, timeout=30, max_idle=300, # 闲置 5 分钟后丢弃,避免复用已断开的连接 kwargs={ "autocommit": True, "prepare_threshold": 0, }, ) await pool.open() checkpointer = AsyncPostgresSaver(pool) await checkpointer.setup() logger.info("Checkpointer 初始化成功") return pool, checkpointer except Exception as e: logger.error(f"Checkpointer 初始化失败 (尝试 {attempt + 1}/{max_retries}): {e}") if pool is not None: try: await pool.close() except Exception: pass if attempt < max_retries - 1: logger.info(f"将在 {retry_delay} 秒后重试...") await asyncio.sleep(retry_delay) retry_delay *= 2 else: logger.error("Checkpointer 初始化失败,已达到最大重试次数") raise async def _discard_psycopg_pool() -> None: """关闭并清除 psycopg 连接池与 Checkpointer,下次调用 get_checkpointer 时会重建。""" global _psycopg_pool, _checkpointer if _psycopg_pool is not None: try: await _psycopg_pool.close() except Exception: pass _psycopg_pool = None _checkpointer = None logger.info("psycopg 连接池与 Checkpointer 已重置") def is_db_connection_error(exc: BaseException) -> bool: """判断异常是否由数据库连接断开引起。""" if isinstance( exc, ( asyncpg.PostgresConnectionError, asyncpg.TooManyConnectionsError, asyncpg.InterfaceError, psycopg.OperationalError, psycopg.InterfaceError, OSError, ), ): return True cause = exc.__cause__ return isinstance(cause, BaseException) and is_db_connection_error(cause) async def reset_db_pools_on_connection_error(exc: BaseException) -> None: """连接异常时重置相关连接池,使后续请求可自动恢复。""" if not is_db_connection_error(exc): return logger.error(f"检测到数据库连接异常,正在重置连接池: {exc}") await _discard_asyncpg_pool() await _discard_psycopg_pool() async def get_checkpointer() -> AsyncPostgresSaver: """ 获取或创建 LangGraph Checkpointer 使用 psycopg AsyncConnectionPool,用于 LangGraph 的状态持久化。 若连接池已失效,会自动丢弃并重建,无需重启服务。 """ global _psycopg_pool, _checkpointer if _checkpointer is None: logger.info("初始化 psycopg 连接池和 Checkpointer...") _psycopg_pool, _checkpointer = await _create_psycopg_checkpointer() return _checkpointer async def close_db_pool(): """关闭所有数据库连接池""" global _asyncpg_pool, _psycopg_pool, _checkpointer if _asyncpg_pool is not None: logger.info("关闭 asyncpg 数据库连接池...") await _asyncpg_pool.close() _asyncpg_pool = None reset_graph_metadata() logger.info("asyncpg 数据库连接池已关闭") if _psycopg_pool is not None: logger.info("关闭 psycopg 连接池...") await _psycopg_pool.close() _psycopg_pool = None _checkpointer = None logger.info("psycopg 连接池已关闭") async def get_db_connection(): """ 获取数据库连接(用于 FastAPI 依赖注入)。 若 acquire 因连接断开而失败,会自动重置连接池, 下一次请求将重新建立连接,无需重启服务。 本次请求仍会返回 500,但不会永久卡死。 """ try: pool = await get_db_pool() async with pool.acquire() as connection: yield connection except ( asyncpg.PostgresConnectionError, asyncpg.TooManyConnectionsError, asyncpg.InterfaceError, OSError, ) as e: await reset_db_pools_on_connection_error(e) raise