chat_graph.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187
  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 app.rag.embeddings import get_embeddings
  50. embeddings = get_embeddings()
  51. memory_mgr = MemoryManager(embeddings)
  52. llm_with_tools = llm.bind_tools([
  53. search_product_by_keyword,
  54. search_article_by_keyword,
  55. search_activity_by_keyword,
  56. ])
  57. builder = StateGraph(ChatState)
  58. # ── 节点 ──
  59. async def classify_intent(state: ChatState) -> dict:
  60. intent = await classifier.classify(
  61. state["query"],
  62. context=str(state.get("context", {})),
  63. )
  64. return {"intent": intent}
  65. async def load_context(state: ChatState) -> dict:
  66. ctx = await agent.load_context(state["user_id"], state.get("child_id"))
  67. return {"context": ctx}
  68. async def llm_call(state: ChatState) -> dict:
  69. """核心 LLM 调用 + Tool"""
  70. chat_prompt = await get_prompt("chat_assistant") or DEFAULT_CHAT_PROMPT
  71. messages = [SystemMessage(content=chat_prompt)]
  72. # 注入家庭上下文 + 用户健康现状
  73. ctx = state.get("context", {})
  74. if isinstance(ctx, dict):
  75. ctx_parts = []
  76. if ctx.get("child_id") is not None:
  77. ctx_parts.append(f"child_id: {ctx['child_id']}")
  78. if ctx.get("report_id") is not None:
  79. ctx_parts.append(f"report_id: {ctx['report_id']}")
  80. if ctx.get("family_id") is not None:
  81. ctx_parts.append(f"family_id: {ctx['family_id']}")
  82. if ctx.get("health_status"):
  83. ctx_parts.append(f"用户健康现状: {ctx['health_status']}")
  84. if ctx.get("diet_preferences"):
  85. ctx_parts.append(f"用户饮食偏好: {ctx['diet_preferences']}")
  86. if ctx_parts:
  87. messages.append(SystemMessage(content="\n家庭上下文:\n" + "\n".join(ctx_parts)))
  88. # 注入长期记忆
  89. try:
  90. memories = await memory_mgr.recall(state["user_id"], state["query"])
  91. if memories:
  92. mem_text = "\n".join([f"- {m}" for m in memories])
  93. messages.append(SystemMessage(
  94. content=f"相关历史对话:\n{mem_text}"
  95. ))
  96. except Exception as e:
  97. logger.warning("召回记忆失败: %s", e)
  98. # 用户消息
  99. messages.append(HumanMessage(content=state["query"]))
  100. response = await llm_with_tools.ainvoke(messages)
  101. answer = response.content
  102. # 提取任务
  103. tasks = await agent.extract_tasks(answer, state["user_id"], state["conversation_id"])
  104. # 提取来源
  105. sources = []
  106. if response.response_metadata.get("tool_calls"):
  107. for tc in response.response_metadata["tool_calls"]:
  108. sources.append({
  109. "type": "tool",
  110. "name": tc.get("name", ""),
  111. "input": tc.get("args", {}),
  112. })
  113. return {
  114. "answer": answer,
  115. "tasks": tasks,
  116. "sources": sources,
  117. "messages": [{"role": "user", "content": state["query"]},
  118. {"role": "assistant", "content": answer}],
  119. }
  120. async def save_memory(state: ChatState) -> dict:
  121. """对话后保存到长期记忆"""
  122. try:
  123. if state.get("messages"):
  124. await memory_mgr.save_conversation(
  125. state["user_id"],
  126. state["conversation_id"],
  127. state["messages"],
  128. )
  129. except Exception as e:
  130. logger.warning("保存记忆失败: %s", e)
  131. return {}
  132. # ── 路由 ──
  133. def route_by_intent(state: ChatState) -> Literal["llm_call", END]:
  134. if state["intent"] in (
  135. Intent.RECOMMEND,
  136. Intent.ANALYSIS,
  137. Intent.HEALTH,
  138. ):
  139. # 这些意图需要更专业的 Agent (Phase 3 实现)
  140. # 当前先走通用 LLM
  141. pass
  142. return "llm_call"
  143. # ── 构建图 ──
  144. builder.add_node("classify_intent", classify_intent)
  145. builder.add_node("load_context", load_context)
  146. builder.add_node("llm_call", llm_call)
  147. builder.add_node("save_memory", save_memory)
  148. builder.add_edge(START, "classify_intent")
  149. builder.add_edge("classify_intent", "load_context")
  150. builder.add_conditional_edges("load_context", route_by_intent)
  151. builder.add_edge("llm_call", "save_memory")
  152. builder.add_edge("save_memory", END)
  153. checkpointer = MemorySaver()
  154. graph = builder.compile(checkpointer=checkpointer)
  155. return graph