This commit is contained in:
silk 2026-06-08 23:26:44 +08:00
parent 27bf97d3d6
commit 61772da143
5 changed files with 293 additions and 140 deletions

View File

@ -26,7 +26,7 @@ from core.llm_catalog import (
resolve_to_api_model,
validate_request_can_use_provider,
)
from core.database import get_db_pool, get_checkpointer
from core.database import get_db_pool, get_checkpointer, reset_db_pools_on_connection_error
from core.mcp_client import get_mcp_client
from core.dependencies import get_current_user, get_moderation_service
from core.exceptions import ModerationError
@ -68,7 +68,7 @@ from services.chat_thread_service import (
get_knowledge_graph_tool_flags,
)
from services.chat_message_file_service import ChatMessageFileService
from services.chat_message_service import ChatMessageService # 新增:消息保存服务
from services.chat_message_service import ChatMessageService, normalize_message_content
from utils.helpers import BaseResponse
from logger.logging import get_logger
@ -622,6 +622,7 @@ async def chat_completion(
logger.error(f"[V2] 保存消息到 chat_messages 表失败: {save_err}")
except Exception as e:
await reset_db_pools_on_connection_error(e)
logger.exception(f"聊天接口错误: {e}")
yield json.dumps({"error": str(e)}, ensure_ascii=False)
@ -957,7 +958,7 @@ async def _save_messages_to_chat_messages_table(
continue
msg_type = msg.type
msg_content = getattr(msg, 'content', '') or ''
msg_content = normalize_message_content(getattr(msg, 'content', '') or '')
# 检查单条消息是否已存在
existing = await conn.fetchval(

View File

@ -9,6 +9,7 @@ from typing import Optional
import asyncio
import asyncpg
import psycopg
from psycopg_pool import AsyncConnectionPool
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
@ -24,136 +25,198 @@ _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 数据库连接池
用于一般的数据库 CRUD 操作
获取或创建 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 数据库连接池: {settings.db_user}@{settings.db_host}:{settings.db_port}/{settings.db_name}")
max_retries = 3
retry_delay = 2 # 秒
for attempt in range(max_retries):
try:
_asyncpg_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, # 连接超时 30 秒
server_settings={
'application_name': 'huoyan-enterprise',
'jit': 'off' # 禁用 JIT 以提高稳定性
}
)
# 测试连接
async with _asyncpg_pool.acquire() as _conn:
await _conn.execute("SELECT 1")
await ensure_graph_metadata(_conn)
logger.info("asyncpg 数据库连接池初始化成功")
break
except Exception as e:
logger.error(f"asyncpg 数据库连接池初始化失败 (尝试 {attempt + 1}/{max_retries}): {e}")
if _asyncpg_pool is not None:
try:
await _asyncpg_pool.close()
except:
pass
_asyncpg_pool = None
if attempt < max_retries - 1:
logger.info(f"将在 {retry_delay} 秒后重试...")
await asyncio.sleep(retry_delay)
retry_delay *= 2 # 指数退避
else:
logger.error("数据库连接池初始化失败,已达到最大重试次数")
raise
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...")
max_retries = 3
retry_delay = 2 # 秒
for attempt in range(max_retries):
try:
_psycopg_pool = AsyncConnectionPool(
conninfo=settings.db_uri_psycopg,
max_size=settings.checkpointer_pool_max_size,
open=False,
timeout=30, # 连接超时 30 秒
kwargs={
"autocommit": True,
"prepare_threshold": 0
},
)
await _psycopg_pool.open()
_checkpointer = AsyncPostgresSaver(_psycopg_pool)
await _checkpointer.setup()
logger.info("Checkpointer 初始化成功")
break
except Exception as e:
logger.error(f"Checkpointer 初始化失败 (尝试 {attempt + 1}/{max_retries}): {e}")
if _psycopg_pool is not None:
try:
await _psycopg_pool.close()
except:
pass
_psycopg_pool = None
_checkpointer = None
if attempt < max_retries - 1:
logger.info(f"将在 {retry_delay} 秒后重试...")
await asyncio.sleep(retry_delay)
retry_delay *= 2 # 指数退避
else:
logger.error("Checkpointer 初始化失败,已达到最大重试次数")
raise
_psycopg_pool, _checkpointer = await _create_psycopg_checkpointer()
return _checkpointer
async def close_db_pool():
"""关闭所有数据库连接池"""
global _asyncpg_pool, _psycopg_pool, _checkpointer
# 关闭 asyncpg 连接池
if _asyncpg_pool is not None:
logger.info("关闭 asyncpg 数据库连接池...")
await _asyncpg_pool.close()
_asyncpg_pool = None
reset_graph_metadata()
logger.info("asyncpg 数据库连接池已关闭")
# 关闭 psycopg 连接池
if _psycopg_pool is not None:
logger.info("关闭 psycopg 连接池...")
await _psycopg_pool.close()
@ -163,8 +226,23 @@ async def close_db_pool():
async def get_db_connection():
"""获取数据库连接(用于依赖注入)"""
pool = await get_db_pool()
async with pool.acquire() as connection:
yield 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

View File

@ -4,13 +4,38 @@
用于保存和查询用户原始消息和AI响应替代从 checkpoint 中解析
"""
import json
from typing import List, Dict, Any, Optional
from typing import List, Dict, Any, Optional, Union
import asyncpg
from logger.logging import get_logger
logger = get_logger(__name__)
def normalize_message_content(content: Union[str, list, dict, None]) -> str:
"""将 LangChain 消息 content可能是 str 或 content blocks 列表)规范为可落库的字符串。"""
if content is None:
return ""
if isinstance(content, str):
return content
if isinstance(content, list):
parts: list[str] = []
for block in content:
if isinstance(block, str):
if block:
parts.append(block)
elif isinstance(block, dict):
text = block.get("text") or block.get("content")
if text is not None and str(text).strip():
parts.append(str(text))
elif block is not None:
parts.append(str(block))
return "\n".join(parts)
if isinstance(content, dict):
text = content.get("text") or content.get("content")
return str(text) if text is not None else json.dumps(content, ensure_ascii=False)
return str(content)
class ChatMessageService:
"""聊天消息服务类"""
@ -59,8 +84,8 @@ class ChatMessageService:
checkpoint_id,
message_index,
'user',
content,
injected_content,
normalize_message_content(content),
normalize_message_content(injected_content) if injected_content else None,
has_files,
json.dumps(metadata) if metadata else None
)
@ -112,7 +137,7 @@ class ChatMessageService:
checkpoint_id,
message_index,
'assistant',
content,
normalize_message_content(content),
json.dumps(metadata) if metadata else None
)
@ -166,7 +191,7 @@ class ChatMessageService:
checkpoint_id,
message_index,
'tool',
content,
normalize_message_content(content),
name,
json.dumps(metadata) if metadata else None
)

View File

@ -13,7 +13,7 @@ function generateUUID() {
}
/** LangChain 落库多为块数组 / 结构化 content需压成可读字符串再给模板与 Markdown */
function normalizeMessageContent(val) {
export function normalizeMessageContent(val) {
if (val == null) return ''
if (typeof val === 'string') return val
if (typeof val === 'number' || typeof val === 'boolean') return String(val)
@ -184,6 +184,11 @@ export const useChatStore = defineStore('chat', () => {
}
const { step, content = '', reasoning_content = '', tool_calls = null, metadata = null, messageData = null } = stepData
const normalizedContent = normalizeMessageContent(content)
const uiTypeFromData = messageData?.type
? normalizeStreamMessageUiType(messageData.type)
: null
const isToolStep = uiTypeFromData === 'tool'
// 查找是否已存在该步骤
let stepIndex = lastMessage.steps.findIndex(s => s.step === step)
@ -204,9 +209,13 @@ export const useChatStore = defineStore('chat', () => {
const currentStep = lastMessage.steps[stepIndex]
// 更新步骤内容(增量更新)
if (content) {
currentStep.content = currentStep.content + content
// 更新步骤内容流式增量tool 消息多为整块 content blocks需先 normalize
if (normalizedContent) {
if (isToolStep || currentStep.messageType === 'tool') {
currentStep.content = normalizedContent
} else {
currentStep.content = currentStep.content + normalizedContent
}
}
// 更新思考内容(增量更新)

View File

@ -432,7 +432,7 @@
{{ tc.function?.name || tc.name }}
</div>
<div v-if="tc.function?.arguments || tc.args" class="tool-call-args">
<code>{{ tc.function?.arguments || tc.args }}</code>
<code>{{ formatToolArgs(tc.function?.arguments || tc.args) }}</code>
</div>
</div>
</div>
@ -872,7 +872,7 @@
import { ref, onMounted, onUnmounted, nextTick, watch, computed } from 'vue'
import { useRouter, useRoute } from 'vue-router'
import { useAuthStore } from '../stores/auth'
import { useChatStore } from '../stores/chat'
import { useChatStore, normalizeMessageContent } from '../stores/chat'
import { useKnowledgeBaseStore } from '../stores/knowledgeBase'
import { marked } from 'marked'
import axios from 'axios'
@ -1247,11 +1247,23 @@ function formatJSON(content) {
}
}
// tool call args JSON
function formatToolArgs(raw) {
if (!raw) return ''
const s = typeof raw === 'string' ? raw : JSON.stringify(raw)
try {
return JSON.stringify(JSON.parse(s), null, 2)
} catch {
return s
}
}
// JSON HTML
function renderJSON(content) {
if (!content) return ''
const text = normalizeMessageContent(content)
if (!text) return ''
try {
const formatted = formatJSON(content)
const formatted = formatJSON(text)
// HTML
const escaped = formatted
.replace(/&/g, '&amp;')
@ -1270,7 +1282,7 @@ function renderJSON(content) {
return `<pre class="json-content"><code>${highlighted}</code></pre>`
} catch (e) {
//
const escaped = content
const escaped = text
.replace(/&/g, '&amp;')
.replace(/</g, '&lt;')
.replace(/>/g, '&gt;')
@ -1778,28 +1790,56 @@ async function sendMessage() {
// langgraph_step使
if (langgraphStep !== undefined && langgraphStep !== null) {
//
const content = messageData.content || ''
// str / LangChain content blocks
const content = normalizeMessageContent(messageData.content)
//
const reasoning_content = messageData.additional_kwargs?.reasoning_content || ''
// 使 tool_calls使 tool_call_chunks
// tool_calls
const existingStep = currentMessage.steps?.find(s => s.step === langgraphStep)
const prevToolCalls = existingStep?.tool_calls ? existingStep.tool_calls.map(tc => ({ ...tc, function: { ...tc.function } })) : []
// args tool_call_chunks args
// invalid_tool_calls LangChain JSON tool_call_chunks
const argChunks = (messageData.tool_call_chunks || [])
const frameToolCalls = messageData.tool_calls || []
const hasNewData = frameToolCalls.some(tc => tc.name || tc.id) || argChunks.some(c => c.args || c.name)
let tool_calls = null
if (messageData.tool_calls && messageData.tool_calls.length > 0) {
tool_calls = messageData.tool_calls
} else if (messageData.tool_call_chunks && messageData.tool_call_chunks.length > 0) {
// tool_call_chunks tool_calls
tool_calls = messageData.tool_call_chunks.map(chunk => ({
id: chunk.id || '',
type: chunk.type || 'function',
function: {
name: chunk.name || '',
arguments: chunk.args || ''
},
name: chunk.name,
args: chunk.args
if (hasNewData || prevToolCalls.length > 0) {
// tool_calls
const merged = prevToolCalls.map(tc => ({
...tc,
function: { ...(tc.function || { name: '', arguments: '' }) }
}))
// tool_calls name / id
frameToolCalls.forEach((tc, idx) => {
if (!merged[idx]) {
merged[idx] = { name: '', args: '', id: '', function: { name: '', arguments: '' } }
}
if (tc.name) { merged[idx].name = tc.name; merged[idx].function.name = tc.name }
if (tc.id) { merged[idx].id = tc.id }
})
// args tool_call_chunks invalid_tool_calls
argChunks.forEach(chunk => {
const idx = typeof chunk.index === 'number' ? chunk.index : 0
if (!merged[idx]) {
merged[idx] = { name: chunk.name || '', args: '', id: chunk.id || '', function: { name: chunk.name || '', arguments: '' } }
}
if (chunk.args) {
merged[idx].args = (merged[idx].args || '') + chunk.args
merged[idx].function.arguments = merged[idx].args
}
if (chunk.name && !merged[idx].name) {
merged[idx].name = chunk.name
merged[idx].function.name = chunk.name
}
})
if (merged.length > 0) tool_calls = merged
}
//
@ -1845,7 +1885,7 @@ async function sendMessage() {
})
} else {
// langgraph_step
const content = messageData.content || ''
const content = normalizeMessageContent(messageData.content)
const newContent = currentMessage.content + content
// additional_kwargs