fix(db): 使用asyncio.shield确保数据库连接正确关闭

1. 简化对话状态重置逻辑并移除冗余代码
2. 重构聊天路由的数据库会话管理
This commit is contained in:
Wenjie Zhang 2025-12-19 02:22:29 +08:00
parent be3abc00b5
commit ba2d0d0007
4 changed files with 163 additions and 199 deletions

View File

@ -549,6 +549,8 @@ async def chat_agent(
thread_id = str(uuid.uuid4()) thread_id = str(uuid.uuid4())
logger.warning(f"No thread_id provided, generated new thread_id: {thread_id}") logger.warning(f"No thread_id provided, generated new thread_id: {thread_id}")
try:
async with db_manager.get_async_session_context() as db:
# Initialize conversation manager # Initialize conversation manager
conv_manager = ConversationManager(db) conv_manager = ConversationManager(db)
@ -574,7 +576,6 @@ async def chat_agent(
logger.error(f"Error loading attachments for thread_id={thread_id}: {e}") logger.error(f"Error loading attachments for thread_id={thread_id}: {e}")
input_context["attachments"] = [] input_context["attachments"] = []
try:
full_msg = None full_msg = None
langgraph_config = {"configurable": input_context} langgraph_config = {"configurable": input_context}
async for msg, metadata in agent.stream_messages(messages, input_context=input_context): async for msg, metadata in agent.stream_messages(messages, input_context=input_context):
@ -645,7 +646,8 @@ async def chat_agent(
# 客户端主动中断连接,检查中断并保存已生成的部分内容 # 客户端主动中断连接,检查中断并保存已生成的部分内容
logger.warning(f"Client disconnected, cancelling stream: {e}") logger.warning(f"Client disconnected, cancelling stream: {e}")
# 保存中断消息到数据库 # Run save in a separate task to avoid cancellation
async def save_cleanup():
async with db_manager.get_async_session_context() as new_db: async with db_manager.get_async_session_context() as new_db:
new_conv_manager = ConversationManager(new_db) new_conv_manager = ConversationManager(new_db)
await save_partial_message( await save_partial_message(
@ -656,6 +658,16 @@ async def chat_agent(
error_type="interrupted", error_type="interrupted",
) )
# Create a task and await it, shielding it from cancellation
# ensuring the DB operation completes even if the stream is cancelled
cleanup_task = asyncio.create_task(save_cleanup())
try:
await asyncio.shield(cleanup_task)
except asyncio.CancelledError:
pass
except Exception as exc:
logger.error(f"Error during cleanup save: {exc}")
# 通知前端中断(可能发送不到,但用于一致性) # 通知前端中断(可能发送不到,但用于一致性)
yield make_chunk(status="interrupted", message="对话已中断", meta=meta) yield make_chunk(status="interrupted", message="对话已中断", meta=meta)
@ -780,6 +792,7 @@ async def resume_agent_chat(
) )
try: try:
async with db_manager.get_async_session_context() as db:
async for msg, metadata in stream_source: async for msg, metadata in stream_source:
# 确保msg有正确的ID结构 # 确保msg有正确的ID结构
msg_dict = msg.model_dump() msg_dict = msg.model_dump()

View File

@ -1,3 +1,4 @@
import asyncio
import json import json
import os import os
import pathlib import pathlib
@ -128,7 +129,9 @@ class DBManager(metaclass=SingletonMeta):
logger.error(f"Async database operation failed: {e}") logger.error(f"Async database operation failed: {e}")
raise raise
finally: finally:
await session.close() # Shield close operation to ensure connection is properly closed even if task is cancelled
# This prevents aiosqlite from raising errors during cancellation
await asyncio.shield(session.close())
def check_first_run(self): def check_first_run(self):
"""检查是否首次运行(同步版本)""" """检查是否首次运行(同步版本)"""

View File

@ -233,10 +233,8 @@
<script setup> <script setup>
import { ref, reactive, onMounted, watch, nextTick, computed, onUnmounted } from 'vue'; import { ref, reactive, onMounted, watch, nextTick, computed, onUnmounted } from 'vue';
import { LoadingOutlined } from '@ant-design/icons-vue';
import { message } from 'ant-design-vue'; import { message } from 'ant-design-vue';
import MessageInputComponent from '@/components/MessageInputComponent.vue' import MessageInputComponent from '@/components/MessageInputComponent.vue'
import AttachmentInputPanel from '@/components/AttachmentInputPanel.vue'
import AttachmentOptionsComponent from '@/components/AttachmentOptionsComponent.vue' import AttachmentOptionsComponent from '@/components/AttachmentOptionsComponent.vue'
import AttachmentStatusIndicator from '@/components/AttachmentStatusIndicator.vue' import AttachmentStatusIndicator from '@/components/AttachmentStatusIndicator.vue'
import AgentMessageComponent from '@/components/AgentMessageComponent.vue' import AgentMessageComponent from '@/components/AgentMessageComponent.vue'
@ -431,27 +429,16 @@ const conversations = computed(() => {
const threadState = currentThreadState.value; const threadState = currentThreadState.value;
// 线 // 线
if (onGoingConvMessages.value.length > 0 && threadState?.isStreaming) { if (onGoingConvMessages.value.length > 0) {
const onGoingConv = { const onGoingConv = {
messages: onGoingConvMessages.value, messages: onGoingConvMessages.value,
status: 'streaming' status: 'streaming'
}; };
return [...historyConvs, onGoingConv]; return [...historyConvs, onGoingConv];
} }
// 使
if (historyConvs.length === 0 && onGoingConvMessages.value.length > 0 && !threadState?.isStreaming) {
const finalConv = {
messages: onGoingConvMessages.value,
status: 'finished'
};
return [finalConv];
}
return historyConvs; return historyConvs;
}); });
const isLoadingThreads = computed(() => chatUIStore.isLoadingThreads);
const isLoadingMessages = computed(() => chatUIStore.isLoadingMessages); const isLoadingMessages = computed(() => chatUIStore.isLoadingMessages);
const isStreaming = computed(() => { const isStreaming = computed(() => {
const threadState = currentThreadState.value; const threadState = currentThreadState.value;
@ -522,55 +509,29 @@ const cleanupThreadState = (threadId) => {
}; };
// ==================== STREAM HANDLING LOGIC ==================== // ==================== STREAM HANDLING LOGIC ====================
const resetOnGoingConv = (threadId = null, preserveMessages = false) => { const resetOnGoingConv = (threadId = null) => {
console.log('🔄 [RESET] Resetting on going conversation:', threadId, preserveMessages); console.log(`🔄 [RESET] Resetting on going conversation: ${new Date().toLocaleTimeString()}.${new Date().getMilliseconds()}`, threadId);
if (threadId) {
// 线 const targetThreadId = threadId || currentChatId.value;
const threadState = getThreadState(threadId);
if (threadState) {
if (threadState.streamAbortController) {
threadState.streamAbortController.abort();
threadState.streamAbortController = null;
}
//
if (preserveMessages) {
//
setTimeout(() => {
if (threadState.onGoingConv) {
threadState.onGoingConv = createOnGoingConvState();
}
}, 100);
} else {
threadState.onGoingConv = createOnGoingConvState();
}
}
} else {
// 线线
const targetThreadId = currentChatId.value;
if (targetThreadId) { if (targetThreadId) {
// 线
const threadState = getThreadState(targetThreadId); const threadState = getThreadState(targetThreadId);
if (threadState) { if (threadState) {
if (threadState.streamAbortController) { if (threadState.streamAbortController) {
threadState.streamAbortController.abort(); threadState.streamAbortController.abort();
threadState.streamAbortController = null; threadState.streamAbortController = null;
} }
if (preserveMessages) {
setTimeout(() => { //
if (threadState.onGoingConv) {
threadState.onGoingConv = createOnGoingConvState(); threadState.onGoingConv = createOnGoingConvState();
} }
}, 100);
} else {
threadState.onGoingConv = createOnGoingConvState();
}
}
} else { } else {
// 线线 // 线线
Object.keys(chatState.threadStates).forEach(tid => { Object.keys(chatState.threadStates).forEach(tid => {
cleanupThreadState(tid); cleanupThreadState(tid);
}); });
} }
}
}; };
const _processStreamChunk = (chunk, threadId) => { const _processStreamChunk = (chunk, threadId) => {
@ -604,10 +565,6 @@ const _processStreamChunk = (chunk, threadId) => {
threadState.streamAbortController = null; threadState.streamAbortController = null;
} }
} }
// Reload messages to show any partial content saved by the backend
fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId, delay: 500 });
resetOnGoingConv(threadId);
return true; return true;
case 'human_approval_required': case 'human_approval_required':
// 使 composable // 使 composable
@ -627,22 +584,17 @@ const _processStreamChunk = (chunk, threadId) => {
if (threadState) { if (threadState) {
threadState.isStreaming = false; threadState.isStreaming = false;
if ((supportsTodo.value || supportsFiles.value) && threadState.agentState) { if ((supportsTodo.value || supportsFiles.value) && threadState.agentState) {
console.log('[AgentState|Final]', { console.log(`[AgentState|Final] ${new Date().toLocaleTimeString()}.${new Date().getMilliseconds()}`, {
threadId, threadId,
todos: threadState.agentState?.todos || [], todos: threadState.agentState?.todos || [],
files: threadState.agentState?.files || [] files: threadState.agentState?.files || []
}); });
} }
} }
//
fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId, delay: 500 })
.finally(() => {
//
resetOnGoingConv(threadId, true);
});
return true; return true;
case 'interrupted': case 'interrupted':
// //
console.warn("[Interrupted] case");
if (threadState) { if (threadState) {
threadState.isStreaming = false; threadState.isStreaming = false;
} }
@ -650,10 +602,6 @@ const _processStreamChunk = (chunk, threadId) => {
if (chunkMessage) { if (chunkMessage) {
message.info(chunkMessage); message.info(chunkMessage);
} }
fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId, delay: 1000 })
.finally(() => {
resetOnGoingConv(threadId, true);
});
return true; return true;
} }
@ -761,7 +709,7 @@ const fetchThreadMessages = async ({ agentId, threadId, delay = 0 }) => {
try { try {
const response = await agentApi.getAgentHistory(agentId, threadId); const response = await agentApi.getAgentHistory(agentId, threadId);
console.log('🔄 [FETCH] Thread messages:', response); console.log(`🔄 [FETCH] Thread messages: ${new Date().toLocaleTimeString()}.${new Date().getMilliseconds()}`, response);
threadMessages.value[threadId] = response.history || []; threadMessages.value[threadId] = response.history || [];
} catch (error) { } catch (error) {
handleChatError(error, 'load'); handleChatError(error, 'load');
@ -1089,12 +1037,21 @@ const handleSendMessage = async () => {
} }
} catch (error) { } catch (error) {
if (error.name !== 'AbortError') { if (error.name !== 'AbortError') {
console.error('Stream error:', error);
handleChatError(error, 'send'); handleChatError(error, 'send');
} else {
console.warn("[Interrupted] Catch");
} }
} finally {
threadState.isStreaming = false; threadState.isStreaming = false;
} finally {
threadState.streamAbortController = null; threadState.streamAbortController = null;
//
fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId, delay: 500 })
.finally(() => {
//
resetOnGoingConv(threadId); resetOnGoingConv(threadId);
scrollController.scrollToBottom();
});
} }
}; };
@ -1208,6 +1165,14 @@ const handleApprovalWithStream = async (approved) => {
threadState.isStreaming = false; threadState.isStreaming = false;
threadState.streamAbortController = null; threadState.streamAbortController = null;
} }
//
fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId, delay: 500 })
.finally(() => {
//
resetOnGoingConv(threadId);
scrollController.scrollToBottom();
});
} }
}; };

View File

@ -153,23 +153,6 @@ const getModelName = (msg) => {
} }
return null; return null;
} }
// Load existing feedback on mount
onMounted(async () => {
if (msg.value?.id) {
try {
const response = await agentApi.getMessageFeedback(msg.value.id)
if (response.has_feedback) {
feedbackState.hasSubmitted = true
feedbackState.rating = response.feedback.rating
feedbackState.reason = response.feedback.reason
}
} catch (error) {
console.error('Failed to load feedback:', error)
}
}
})
// Handle like action // Handle like action
const likeThisResponse = async (msg) => { const likeThisResponse = async (msg) => {
if (feedbackState.hasSubmitted) { if (feedbackState.hasSubmitted) {