fix agent

This commit is contained in:
jxxghp
2025-11-01 10:39:08 +08:00
parent d523c7c916
commit 438d3210bc
18 changed files with 145 additions and 71 deletions
+43 -16
View File
@@ -15,9 +15,15 @@ from langchain_core.runnables.history import RunnableWithMessageHistory
from app.agent.memory import ConversationMemoryManager
from app.agent.prompt import PromptManager
from app.agent.tools import MoviePilotToolFactory
from app.chain import ChainBase
from app.core.config import settings
from app.helper.message import MessageHelper
from app.log import logger
from app.schemas import Notification
class AgentChain(ChainBase):
pass
class StreamingCallbackHandler(AsyncCallbackHandler):
@@ -51,9 +57,13 @@ class StreamingCallbackHandler(AsyncCallbackHandler):
class MoviePilotAgent:
"""MoviePilot AI智能体"""
def __init__(self, session_id: str, user_id: str = None):
def __init__(self, session_id: str, user_id: str = None,
channel: str = None, source: str = None, username: str = None):
self.session_id = session_id
self.user_id = user_id
self.channel = channel # 消息渠道
self.source = source # 消息来源
self.username = username # 用户名
# 消息助手
self.message_helper = MessageHelper()
@@ -173,14 +183,14 @@ class MoviePilotAgent:
def _initialize_prompt() -> ChatPromptTemplate:
"""初始化提示词模板"""
try:
prompt = ChatPromptTemplate.from_messages([
prompt_template = ChatPromptTemplate.from_messages([
("system", "{system_prompt}"),
MessagesPlaceholder(variable_name="chat_history"),
("user", "{input}"),
MessagesPlaceholder(variable_name="agent_scratchpad"),
])
logger.info("LangChain提示词模板初始化成功")
return prompt
return prompt_template
except Exception as e:
logger.error(f"初始化提示词失败: {e}")
raise e
@@ -236,11 +246,8 @@ class MoviePilotAgent:
# 获取Agent回复
agent_message = await self.callback_handler.get_message()
# 发送Agent回复给用户
self.message_helper.put(
message=agent_message,
role="system"
)
# 发送Agent回复给用户(通过原渠道)
self._send_message_to_channel(agent_message)
# 添加Agent回复到记忆
await self.memory_manager.add_memory(
@@ -255,14 +262,23 @@ class MoviePilotAgent:
except Exception as e:
error_message = f"处理消息时发生错误: {str(e)}"
logger.error(error_message)
# 发送错误消息给用户
self.message_helper.put(
message=error_message,
role="system",
title="MoviePilot助手错误"
)
# 发送错误消息给用户(通过原渠道)
self._send_message_to_channel(error_message)
return error_message
def _send_message_to_channel(self, message: str, title: str = "MoviePilot助手"):
"""通过原渠道发送消息给用户"""
AgentChain().post_message(
Notification(
channel=self.channel,
source=self.source,
userid=self.user_id,
username=self.username,
title=title,
text=message
)
)
async def _execute_agent(self, input_context: Dict[str, Any]) -> Dict[str, Any]:
"""执行LangChain Agent"""
try:
@@ -322,20 +338,31 @@ class AgentManager:
await agent.cleanup()
self.active_agents.clear()
async def process_message(self, session_id: str, user_id: str, message: str) -> str:
async def process_message(self, session_id: str, user_id: str, message: str,
channel: str = None, source: str = None, username: str = None) -> str:
"""处理用户消息"""
# 获取或创建Agent实例
if session_id not in self.active_agents:
logger.info(f"创建新的AI智能体实例,session_id: {session_id}, user_id: {user_id}")
agent = MoviePilotAgent(
session_id=session_id,
user_id=user_id
user_id=user_id,
channel=channel,
source=source,
username=username
)
agent.memory_manager = self.memory_manager
self.active_agents[session_id] = agent
else:
agent = self.active_agents[session_id]
agent.user_id = user_id # 确保user_id是最新的
# 更新渠道信息
if channel:
agent.channel = channel
if source:
agent.source = source
if username:
agent.username = username
# 处理消息
return await agent.process_message(message)