chat_graph.py 6.4 KB

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