chat_graph.py 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178
  1. from typing import TypedDict, Literal
  2. from langgraph.graph import StateGraph, START, END
  3. from langgraph.checkpoint.memory import MemorySaver
  4. from langchain_openai import ChatOpenAI
  5. from langchain_core.messages import SystemMessage, HumanMessage
  6. from app.agents.intent_classifier import IntentClassifier, Intent
  7. from app.agents.chat_agent import ChatAgent
  8. from app.tools.product_tools import (
  9. search_product_by_keyword,
  10. search_article_by_keyword,
  11. search_activity_by_keyword,
  12. )
  13. from app.memory.store import MemoryManager
  14. from app.config import settings
  15. import logging
  16. logger = logging.getLogger(__name__)
  17. class ChatState(TypedDict):
  18. query: str
  19. user_id: int
  20. conversation_id: str
  21. child_id: int | None
  22. intent: Intent | None
  23. context: dict | None
  24. messages: list | None
  25. answer: str | None
  26. tasks: list[dict]
  27. sources: list[dict]
  28. CHAT_SYSTEM_PROMPT = """你是一个儿童成长家庭助手, 回答关于孩子成长、健康、教育的各种问题。
  29. 你可以使用搜索工具查找商品、活动和文章来辅助回答。
  30. 回答原则:
  31. 1. 用中文, 语气温暖亲切
  32. 2. 如果用户提到具体孩子, 参考提供的家庭上下文
  33. 3. 需要推荐时使用搜索工具
  34. 4. 可以生成 [TASK: {"title": "任务名", "dimension": "身/心/智/行/富", "points": 10}] 标记来创建行动任务
  35. 5. 不要编造医疗建议, 严重问题建议咨询医生
  36. """
  37. def create_chat_graph():
  38. """创建聊天 StateGraph"""
  39. agent = ChatAgent()
  40. classifier = IntentClassifier()
  41. llm = ChatOpenAI(
  42. model=settings.llm_model,
  43. api_key=settings.llm_api_key,
  44. base_url=settings.llm_base_url,
  45. temperature=settings.llm_temperature,
  46. )
  47. # 初始化嵌入模型用于记忆存储
  48. from langchain_openai import OpenAIEmbeddings
  49. embeddings = OpenAIEmbeddings(
  50. model=settings.embedding_model,
  51. api_key=settings.effective_embedding_api_key,
  52. base_url=settings.effective_embedding_base_url,
  53. )
  54. memory_mgr = MemoryManager(embeddings)
  55. llm_with_tools = llm.bind_tools([
  56. search_product_by_keyword,
  57. search_article_by_keyword,
  58. search_activity_by_keyword,
  59. ])
  60. builder = StateGraph(ChatState)
  61. # ── 节点 ──
  62. async def classify_intent(state: ChatState) -> dict:
  63. intent = await classifier.classify(
  64. state["query"],
  65. context=str(state.get("context", {})),
  66. )
  67. return {"intent": intent}
  68. async def load_context(state: ChatState) -> dict:
  69. ctx = await agent.load_context(state["user_id"], state.get("child_id"))
  70. return {"context": ctx}
  71. async def llm_call(state: ChatState) -> dict:
  72. """核心 LLM 调用 + Tool"""
  73. messages = [SystemMessage(content=CHAT_SYSTEM_PROMPT)]
  74. # 注入家庭上下文
  75. ctx = state.get("context", {})
  76. if ctx:
  77. ctx_text = f"\n家庭上下文:\n{ctx}"
  78. messages.append(SystemMessage(content=ctx_text))
  79. # 注入长期记忆
  80. try:
  81. memories = await memory_mgr.recall(state["user_id"], state["query"])
  82. if memories:
  83. mem_text = "\n".join([f"- {m}" for m in memories])
  84. messages.append(SystemMessage(
  85. content=f"相关历史对话:\n{mem_text}"
  86. ))
  87. except Exception as e:
  88. logger.warning("召回记忆失败: %s", e)
  89. # 用户消息
  90. messages.append(HumanMessage(content=state["query"]))
  91. response = await llm_with_tools.ainvoke(messages)
  92. answer = response.content
  93. # 提取任务
  94. tasks = await agent.extract_tasks(answer, state["user_id"], state["conversation_id"])
  95. # 提取来源
  96. sources = []
  97. if response.response_metadata.get("tool_calls"):
  98. for tc in response.response_metadata["tool_calls"]:
  99. sources.append({
  100. "type": "tool",
  101. "name": tc.get("name", ""),
  102. "input": tc.get("args", {}),
  103. })
  104. return {
  105. "answer": answer,
  106. "tasks": tasks,
  107. "sources": sources,
  108. "messages": [{"role": "user", "content": state["query"]},
  109. {"role": "assistant", "content": answer}],
  110. }
  111. async def save_memory(state: ChatState) -> dict:
  112. """对话后保存到长期记忆"""
  113. try:
  114. if state.get("messages"):
  115. await memory_mgr.save_conversation(
  116. state["user_id"],
  117. state["conversation_id"],
  118. state["messages"],
  119. )
  120. except Exception as e:
  121. logger.warning("保存记忆失败: %s", e)
  122. return {}
  123. # ── 路由 ──
  124. def route_by_intent(state: ChatState) -> Literal["llm_call", END]:
  125. if state["intent"] in (
  126. Intent.RECOMMEND,
  127. Intent.ANALYSIS,
  128. Intent.HEALTH,
  129. ):
  130. # 这些意图需要更专业的 Agent (Phase 3 实现)
  131. # 当前先走通用 LLM
  132. pass
  133. return "llm_call"
  134. # ── 构建图 ──
  135. builder.add_node("classify_intent", classify_intent)
  136. builder.add_node("load_context", load_context)
  137. builder.add_node("llm_call", llm_call)
  138. builder.add_node("save_memory", save_memory)
  139. builder.add_edge(START, "classify_intent")
  140. builder.add_edge("classify_intent", "load_context")
  141. builder.add_conditional_edges("load_context", route_by_intent)
  142. builder.add_edge("llm_call", "save_memory")
  143. builder.add_edge("save_memory", END)
  144. checkpointer = MemorySaver()
  145. graph = builder.compile(checkpointer=checkpointer)
  146. return graph