fix(chat_router): 修复流式消息处理中未正确累积内容的问题 Question: 上下文比较长时,回答速度很慢,CPU只吃了单核,一直占满资源
Fixes #410 确保在流式处理过程中正确累积消息内容,并在各种中断情况下(如敏感内容检测、错误处理等)保存已累积的内容。同时添加对用户会话的权限检查
This commit is contained in:
parent
26b25df027
commit
606ff9f122
@ -5,7 +5,7 @@ import uuid
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Query, UploadFile, File
|
||||
from fastapi.responses import StreamingResponse
|
||||
from langchain.messages import AIMessageChunk, HumanMessage
|
||||
from langchain.messages import AIMessageChunk, HumanMessage, AIMessage
|
||||
from langgraph.types import Command
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@ -577,13 +577,16 @@ async def chat_agent(
|
||||
input_context["attachments"] = []
|
||||
|
||||
full_msg = None
|
||||
accumulated_content = []
|
||||
langgraph_config = {"configurable": input_context}
|
||||
async for msg, metadata in agent.stream_messages(messages, input_context=input_context):
|
||||
if isinstance(msg, AIMessageChunk):
|
||||
full_msg = msg if not full_msg else full_msg + msg
|
||||
content_for_check = full_msg.content[-20:]
|
||||
accumulated_content.append(msg.content)
|
||||
|
||||
content_for_check = "".join(accumulated_content[-10:])
|
||||
if conf.enable_content_guard and await content_guard.check_with_keywords(content_for_check):
|
||||
logger.warning("Sensitive content detected in stream")
|
||||
full_msg = AIMessage(content="".join(accumulated_content))
|
||||
await save_partial_message(conv_manager, thread_id, full_msg, "content_guard_blocked")
|
||||
meta["time_cost"] = asyncio.get_event_loop().time() - start_time
|
||||
yield make_chunk(status="interrupted", message="检测到敏感内容,已中断输出", meta=meta)
|
||||
@ -606,6 +609,9 @@ async def chat_agent(
|
||||
logger.error(f"Error processing tool message: {e}")
|
||||
pass
|
||||
|
||||
if not full_msg and accumulated_content:
|
||||
full_msg = AIMessage(content="".join(accumulated_content))
|
||||
|
||||
if (
|
||||
conf.enable_content_guard
|
||||
and hasattr(full_msg, "content")
|
||||
@ -650,6 +656,10 @@ async def chat_agent(
|
||||
|
||||
# Run save in a separate task to avoid cancellation
|
||||
async def save_cleanup():
|
||||
nonlocal full_msg
|
||||
if not full_msg and accumulated_content:
|
||||
full_msg = AIMessage(content="".join(accumulated_content))
|
||||
|
||||
async with db_manager.get_async_session_context() as new_db:
|
||||
new_conv_manager = ConversationManager(new_db)
|
||||
await save_partial_message(
|
||||
@ -679,6 +689,9 @@ async def chat_agent(
|
||||
error_msg = f"Error streaming messages: {e}"
|
||||
error_type = "unexpected_error"
|
||||
|
||||
if not full_msg and accumulated_content:
|
||||
full_msg = AIMessage(content="".join(accumulated_content))
|
||||
|
||||
# 保存错误消息到数据库
|
||||
async with db_manager.get_async_session_context() as new_db:
|
||||
new_conv_manager = ConversationManager(new_db)
|
||||
@ -888,6 +901,7 @@ async def get_agent_history(
|
||||
|
||||
# Use new storage system ONLY
|
||||
conv_manager = ConversationManager(db)
|
||||
await _require_user_conversation(conv_manager, thread_id, str(current_user.id))
|
||||
messages = await conv_manager.get_messages_by_thread_id(thread_id)
|
||||
|
||||
# 当前用户ID - 用于过滤反馈
|
||||
|
||||
Loading…
Reference in New Issue
Block a user