Files
life-echo/api/agents/conversation_agent.py

121 lines
4.2 KiB
Python
Raw Normal View History

2026-01-07 11:56:53 +08:00
"""
对话 Agent基于访谈问题清单动态选择问题实时生成回应
"""
from typing import List, Optional
from langchain.chains import ConversationChain
from langchain.memory import ConversationBufferMemory
from langchain.prompts import PromptTemplate
from services.llm_service import llm_service
2026-01-07 11:56:53 +08:00
from .prompts import ConversationStage, get_conversation_prompt
class ConversationAgent:
"""对话 Agent"""
def __init__(self):
# 使用 LLM 服务获取 LLM 实例
self.llm = llm_service.get_llm()
2026-01-07 11:56:53 +08:00
# 对话记忆
self.memories: dict[str, ConversationBufferMemory] = {}
def _get_memory(self, conversation_id: str) -> ConversationBufferMemory:
"""获取或创建对话记忆"""
if conversation_id not in self.memories:
self.memories[conversation_id] = ConversationBufferMemory(
return_messages=True,
memory_key="history"
)
return self.memories[conversation_id]
def generate_response(
self,
conversation_id: str,
user_message: str,
current_stage: Optional[ConversationStage] = None,
covered_topics: Optional[List[str]] = None
) -> str:
"""
生成 Agent 回应
Args:
conversation_id: 对话 ID
user_message: 用户消息
current_stage: 当前对话阶段
covered_topics: 已聊过的话题列表
Returns:
Agent 回应文本
"""
if current_stage is None:
current_stage = ConversationStage.CHILDHOOD
if covered_topics is None:
covered_topics = []
# 如果没有配置 LLM返回默认回应
if not self.llm:
return "抱歉LLM 服务未配置。请设置 DEEPSEEK_API_KEY 或 LLM_API_KEY 环境变量。"
2026-01-07 11:56:53 +08:00
# 获取系统提示词
system_prompt = get_conversation_prompt(current_stage, covered_topics, user_message)
# 获取对话记忆
memory = self._get_memory(conversation_id)
# 创建对话链
prompt_template = PromptTemplate(
input_variables=["history", "input"],
template=f"{system_prompt}\n\n{{history}}\n\nHuman: {{input}}\n\nAssistant:"
)
chain = ConversationChain(
llm=self.llm,
prompt=prompt_template,
memory=memory,
verbose=False
)
# 生成回应
response = chain.predict(input=user_message)
return response
def detect_stage(self, conversation_id: str, user_message: str) -> ConversationStage:
"""
检测对话阶段
Args:
conversation_id: 对话 ID
user_message: 用户消息
Returns:
检测到的对话阶段
"""
# 简单的关键词检测(实际应该使用更智能的方法)
message_lower = user_message.lower()
if any(word in message_lower for word in ["童年", "小时候", "出生", "家庭背景"]):
return ConversationStage.CHILDHOOD
elif any(word in message_lower for word in ["上学", "学校", "老师", "同学", "教育"]):
return ConversationStage.EDUCATION
elif any(word in message_lower for word in ["工作", "职业", "事业", "公司", "同事"]):
return ConversationStage.CAREER
elif any(word in message_lower for word in ["伴侣", "孩子", "家庭", "家人", "结婚"]):
return ConversationStage.FAMILY
elif any(word in message_lower for word in ["信念", "价值观", "座右铭", "坚持", "原则"]):
return ConversationStage.BELIEFS
elif any(word in message_lower for word in ["总结", "回顾", "感激", "希望", "未来"]):
return ConversationStage.SUMMARY
else:
# 默认返回当前阶段或童年阶段
return ConversationStage.CHILDHOOD
def clear_memory(self, conversation_id: str):
"""清除对话记忆"""
if conversation_id in self.memories:
del self.memories[conversation_id]