huoyan-enterprise/backend/core/database.py

249 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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