249 lines
8.3 KiB
Python
249 lines
8.3 KiB
Python
"""
|
||
数据库连接管理模块
|
||
|
||
统一管理所有数据库连接池:
|
||
- 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
|
||
|