chat_graph.py 7.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219
  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. from app.portrait_service import get_user_portrait_text
  17. from app.tools.java_client import JavaClient
  18. import logging
  19. logger = logging.getLogger(__name__)
  20. class ChatState(TypedDict):
  21. query: str
  22. user_id: int
  23. conversation_id: str
  24. child_id: int | None
  25. intent: Intent | None
  26. context: dict | None
  27. self_check_result: str | None # P1-2: 自检结果 JSON
  28. messages: list | None
  29. answer: str | None
  30. tasks: list[dict]
  31. sources: list[dict]
  32. DEFAULT_CHAT_PROMPT = """你是一个儿童成长家庭助手, 回答关于孩子成长、健康、教育的各种问题。
  33. 你可以使用搜索工具查找商品、活动和文章来辅助回答。
  34. 回答原则:
  35. 1. 用中文, 语气温暖亲切
  36. 2. 如果用户提到具体孩子, 参考提供的家庭上下文
  37. 3. 需要推荐时使用搜索工具
  38. 4. 可以生成 [TASK: {"title": "任务名", "dimension": "身/心/智/行/富", "points": 10}] 标记来创建行动任务
  39. 5. 不要编造医疗建议, 严重问题建议咨询医生
  40. """
  41. def create_chat_graph():
  42. """创建聊天 StateGraph"""
  43. agent = ChatAgent()
  44. classifier = IntentClassifier()
  45. llm = ChatOpenAI(
  46. model=settings.llm_model,
  47. api_key=settings.llm_api_key,
  48. base_url=settings.llm_base_url,
  49. temperature=settings.llm_temperature,
  50. )
  51. # 初始化嵌入模型用于记忆存储
  52. from app.rag.embeddings import get_embeddings
  53. embeddings = get_embeddings()
  54. memory_mgr = MemoryManager(embeddings)
  55. java = JavaClient()
  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. member_id = state.get("child_id")
  78. portrait_text = await get_user_portrait_text(java, state["user_id"], member_id)
  79. if portrait_text:
  80. messages.insert(1, SystemMessage(content=portrait_text))
  81. # P1-2: 自检结果注入(如有)
  82. self_check_result = state.get("self_check_result")
  83. if self_check_result:
  84. try:
  85. import json as _json
  86. sc = _json.loads(self_check_result)
  87. sc_text = f"用户最近一次五维自检结果:\n总分 {sc.get('totalScore', 0)}/45,"
  88. dims = sc.get('dimensions') or sc.get('scores') or {}
  89. if isinstance(dims, list):
  90. dim_lines = [f"{d.get('name','')}: {d.get('score',0)}分" for d in dims]
  91. elif isinstance(dims, dict):
  92. dim_lines = [f"{k}: {v}分" for k, v in dims.items()]
  93. else:
  94. dim_lines = [str(dims)]
  95. sc_text += " | ".join(dim_lines)
  96. insight = sc.get('familyInsight') or sc.get('aiInsight')
  97. if insight:
  98. sc_text += f"\nAI 洞察: {insight}"
  99. messages.insert(1, SystemMessage(content=sc_text))
  100. except Exception as _e:
  101. logger.warning("解析 self_check_result 失败: %s", _e)
  102. # 注入家庭上下文 + 用户健康现状
  103. ctx = state.get("context", {})
  104. if isinstance(ctx, dict):
  105. ctx_parts = []
  106. if ctx.get("child_id") is not None:
  107. ctx_parts.append(f"child_id: {ctx['child_id']}")
  108. if ctx.get("report_id") is not None:
  109. ctx_parts.append(f"report_id: {ctx['report_id']}")
  110. if ctx.get("family_id") is not None:
  111. ctx_parts.append(f"family_id: {ctx['family_id']}")
  112. if ctx.get("health_status"):
  113. ctx_parts.append(f"用户健康现状: {ctx['health_status']}")
  114. if ctx.get("diet_preferences"):
  115. ctx_parts.append(f"用户饮食偏好: {ctx['diet_preferences']}")
  116. if ctx_parts:
  117. messages.append(SystemMessage(content="\n家庭上下文:\n" + "\n".join(ctx_parts)))
  118. # 注入长期记忆
  119. try:
  120. memories = await memory_mgr.recall(state["user_id"], state["query"])
  121. if memories:
  122. mem_text = "\n".join([f"- {m}" for m in memories])
  123. messages.append(SystemMessage(
  124. content=f"相关历史对话:\n{mem_text}"
  125. ))
  126. except Exception as e:
  127. logger.warning("召回记忆失败: %s", e)
  128. # 用户消息
  129. messages.append(HumanMessage(content=state["query"]))
  130. response = await llm_with_tools.ainvoke(messages)
  131. answer = response.content
  132. # 提取任务
  133. tasks = await agent.extract_tasks(answer, state["user_id"], state["conversation_id"])
  134. # 提取来源
  135. sources = []
  136. if response.response_metadata.get("tool_calls"):
  137. for tc in response.response_metadata["tool_calls"]:
  138. sources.append({
  139. "type": "tool",
  140. "name": tc.get("name", ""),
  141. "input": tc.get("args", {}),
  142. })
  143. return {
  144. "answer": answer,
  145. "tasks": tasks,
  146. "sources": sources,
  147. "messages": [{"role": "user", "content": state["query"]},
  148. {"role": "assistant", "content": answer}],
  149. }
  150. async def save_memory(state: ChatState) -> dict:
  151. """对话后保存到长期记忆"""
  152. try:
  153. if state.get("messages"):
  154. await memory_mgr.save_conversation(
  155. state["user_id"],
  156. state["conversation_id"],
  157. state["messages"],
  158. )
  159. except Exception as e:
  160. logger.warning("保存记忆失败: %s", e)
  161. return {}
  162. # ── 路由 ──
  163. def route_by_intent(state: ChatState) -> Literal["llm_call", END]:
  164. if state["intent"] in (
  165. Intent.RECOMMEND,
  166. Intent.ANALYSIS,
  167. Intent.HEALTH,
  168. ):
  169. # 这些意图需要更专业的 Agent (Phase 3 实现)
  170. # 当前先走通用 LLM
  171. pass
  172. return "llm_call"
  173. # ── 构建图 ──
  174. builder.add_node("classify_intent", classify_intent)
  175. builder.add_node("load_context", load_context)
  176. builder.add_node("llm_call", llm_call)
  177. builder.add_node("save_memory", save_memory)
  178. builder.add_edge(START, "classify_intent")
  179. builder.add_edge("classify_intent", "load_context")
  180. builder.add_conditional_edges("load_context", route_by_intent)
  181. builder.add_edge("llm_call", "save_memory")
  182. builder.add_edge("save_memory", END)
  183. checkpointer = MemorySaver()
  184. graph = builder.compile(checkpointer=checkpointer)
  185. return graph